17 Commits
Author SHA1 Message Date
Nixevol a0a24b29b0 fix: 统一对话密码锁键为发送方加对方不含 IP 2026-09-30 15:33:15 +08:00
Nixevol c15e8dc4ae merge: admin-web 2026-09-30 15:28:09 +08:00
Nixevol 73408ab96e merge: identity u01-u05 2026-09-30 15:28:06 +08:00
Nixevol 08c0eb09be fix: 自适应表格高度并改用中文标签 2026-09-30 15:22:11 +08:00
Nixevol 2b68b6768c fix: 群详情成员分页并与后端能力对齐 2026-09-30 15:19:38 +08:00
Nixevol 1448e7a9ba fix: 补齐投递记录时间筛选统计与分页 2026-09-30 15:18:13 +08:00
Nixevol 1ec597c3cc fix: 修复端列表解锁、批量选择与编辑提交 2026-09-30 15:16:36 +08:00
Nixevol 5eb4fad15e fix: 校正一次性密码展示与 CSV 下载 2026-09-30 15:14:05 +08:00
Nixevol 898b9d2934 fix: 统一管理后台 401 与加载失败处理 2026-09-30 15:12:16 +08:00
Nixevol eff4bd53be fix: 改密计入锁定并清理过期管理员会话 2026-09-30 15:05:55 +08:00
Nixevol ad7c50d504 fix: 校验端默认延迟不超过调度上限 2026-09-30 15:05:37 +08:00
Nixevol fa62c4c97c fix: 目录搜索转义 LIKE 通配符并拒绝编号 inline 2026-09-30 15:00:05 +08:00
Nixevol 5dac4d87b5 fix: 锁定计数表过期清理并限制总量 2026-09-30 14:59:54 +08:00
Nixevol 4852be91be fix: 开启自助注册须先有 8-64 字符安全码 2026-09-30 14:59:46 +08:00
Nixevol 30330cea2c fix: 批量导入并发哈希并校正 CSV 行号 2026-09-30 14:59:35 +08:00
Nixevol 512b8c2ba1 fix: 端启停改为仅由 identity 执行 2026-09-30 14:56:18 +08:00
Nixevol 5176e04995 fix: 限制管理接口请求体大小与读取时间 2026-09-30 14:55:36 +08:00
66 changed files with 3575 additions and 325 deletions
+136
View File
@@ -648,6 +648,15 @@
- 备选方案:等 C-03 占用 0003 后再用 0004(rebase 时改号)。 - 备选方案:等 C-03 占用 0003 后再用 0004(rebase 时改号)。
- 影响:若 C-03 先合入并占用 0003,本文件 rebase 时改号。 - 影响:若 C-03 先合入并占用 0003,本文件 rebase 时改号。
### 复审修复 U-03
1. **对话密码锁键统一为发送方+对方,不含 IP**
- 原条款:DEVELOPMENT 第 5 节按「发送方 + 对方」计数;锁定期间 unlock、**带密码的**发送和拉人进群返回 `rate_limited`。PRD F15 / D24:改密清按对方计的总数;删除后编号可复用。issue #41。
- 实际做法:`LockTalkPair` 键为 `{Kind, 发送方, 对方}`,IP 留空。unlock / send / 进群共用该键;没带密码直接 `talk_password_required`,不计次、不因已锁改成 `rate_limited`。`UnlockTalk`/`CheckTalkPasswordForJoin` 仍保留 `remoteIP` 参数以免改接线签名。后台 `PUT talk-password` 有 Identity 时调 `SelfSetTalkPassword`(已清 `LockTalkTarget`),无 Identity 时改库后 `Clear(LockTalkTarget)`。删除端调用 `LoginLocks.ClearAllForEndpoint`。不改 group `emit`,不修 H-02。
- 原因:原先 unlock 带 IP、send 成功清零不带 IP、群加人用空 IP,同一发送方换 IP 可再猜 10 次;没带密码也会被已锁挡成 `rate_limited`,SDK 会退避最多 1 小时。
- 备选方案:对话密码也按编号+IP(否决,与第 5 节原文及 F15 不一致)。
- 影响:同一对端合计 10 次错即锁;没带密码始终是 `talk_password_required`;后台改密可解除按对方暂停;删端后同编号重开不继承锁定。
## 后台接口 A ## 后台接口 A
### A1 2026-09-30 ### A1 2026-09-30
@@ -760,6 +769,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 ## 后台网页 W
1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`**
@@ -818,6 +872,61 @@
- 备选:仅文档要求手跑 e2e;未采用。 - 备选:仅文档要求手跑 e2e;未采用。
- 影响:`task itest` 变长;需本机 Go/Node/Chromium。 - 影响:`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) ### S1.1 传输层可注入假实现(Go / JS)
- 相关文档:DEVELOPMENT 第 9 节单元测试要求「用假的 MQTT/HTTP,不要起真实服务器」。 - 相关文档:DEVELOPMENT 第 9 节单元测试要求「用假的 MQTT/HTTP,不要起真实服务器」。
@@ -1237,3 +1346,30 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
5. **建议的正确方向** 5. **建议的正确方向**
- 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。 - 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。
- 覆盖 presence 等其他同步 `PublishDown`,而不只包一层 `emit`。 - 覆盖 presence 等其他同步 `PublishDown`,而不只包一层 `emit`。
### 复审修复 U-01
1. **开启自助注册必须已有 8–64 字符安全码**
- 原条款:PRD F23 / D14:管理员开启并设置注册安全码(8–64 字符);注册必须带当前安全码。
- 实际做法:`register()` 在锁定检查前,若存储码不是 8–64 字符则 403 `registration_closed` 并记不含码的警告,不计入锁定。`PUT /api/admin/registration` 在同一写操作末尾读回开关与安全码,开启且码不合法则 400 并回滚;允许一次提交 `enabled`+`code`/`generate`。后台无已保存安全码时禁用开关。
- 原因:新库无码行或空码时 `constantTimeEqual("", "")` 为真,只开开关即可裸注册。
- 备选方案:仅拦管理 PUT、不拦已处于「开启+空码」的旧库(否决,缺少纵深防御)。
- 影响:原先「先开开关再设码」的两步会 400;须先设码或一次提交开启与码。
### 复审修复 U-04
1. **锁定计数表过期清理与总量上限**
- 原条款:DEVELOPMENT 第 5 节锁定计数;issue #42。
- 实际做法:`internal/auth/locks.go` 的 Fail/Check 顺手删除已过期且最近失败在窗口外的条目;每 1024 次或每分钟全表扫描;默认上限 65536,超出时优先淘汰最旧的非锁定条目。生产接线使用 `auth.NewLoginLocks()`。
- 原因:注册安全码错误与错误 API 令牌按 IP 建条目,轮换地址会使 map 只增不减。
- 备选方案:一并改 `internal/admin/memlock.go`(否决,本波只改 locks.go;admin 测试用内存锁若仍独立注入需后续对齐)。未改对话密码锁键(U-03)。
- 影响:过期未锁定条目会被回收;极端并发失败时最早的非锁定计数可能被挤出。
### 复审修复 U-05
1. **目录 LIKE 转义;注册拒绝编号 inline**
- 原条款:DEVELOPMENT 6.5 按编号前缀或名称包含匹配;issue #43。
- 实际做法:目录查询对 `\`、`%`、`_` 转义并 `ESCAPE '\'`。注册拒绝编号 `inline`(与 mochi 内联客户端 ClientID 同名)。
- 原因:未转义时搜 `e_ab` 会命中 `exab…`,搜 `%` 返回全部端。
- 备选方案:在 `internal/protocol.ValidEndpointID` 加保留字,使开通/导入一并拒绝(否决,本波不改 protocol 与 `internal/admin/endpoints.go`)。后台开通与批量导入仍可能使用 `inline`,留给后续波次。
- 影响:目录下划线按字面匹配;自助注册不能占用 `inline`。
+3
View File
@@ -396,6 +396,8 @@ Authorization: Bearer nxm_...
``` ```
- `generate: true` 时服务器生成 16 位安全码并忽略请求里的 `code`。 - `generate: true` 时服务器生成 16 位安全码并忽略请求里的 `code`。
- 允许一次提交 `{"enabled":true,"code":"..."}` 或 `{"enabled":true,"generate":true}`。
- 最终状态为开启时,安全码必须是 8–64 字符;仅 `{"enabled":true}` 且库中没有合法安全码时返回 400,整次更新回滚。
- 成功 `data` 同 GET。 - 成功 `data` 同 GET。
--- ---
@@ -557,6 +559,7 @@ Authorization: Bearer nxm_...
"name": "一组", "name": "一组",
"owner_id": "a", "owner_id": "a",
"created_at_ms": 1750000000000, "created_at_ms": 1750000000000,
"member_total": 1,
"members": [ "members": [
{"id": "a", "name": "", "online": true, "joined_at_ms": 1750000000000} {"id": "a", "name": "", "online": true, "joined_at_ms": 1750000000000}
], ],
+87
View File
@@ -134,6 +134,66 @@ func TestRegistrationToggleAndGenerate(t *testing.T) {
} }
} }
func TestRegistrationEnableRequiresCode(t *testing.T) {
db, srv, client := setupA3(t)
base := srv.URL
res := doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true}`, csrf())
env := decodeEnv(t, res)
if res.StatusCode != 400 || env.OK {
t.Fatalf("enable without code: %d %+v", res.StatusCode, env)
}
if env.Error == nil || !strings.Contains(env.Error.Message, "8–64") {
t.Fatalf("want 8–64 message, got %+v", env.Error)
}
res = doReq(t, client, http.MethodGet, base+"/api/admin/registration", "", nil)
env = decodeEnv(t, res)
var got map[string]any
_ = json.Unmarshal(env.Data, &got)
if got["enabled"] != false {
t.Fatalf("switch must stay off after 400: %v", got)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM settings WHERE key = ? AND value = '1'`, "registration_enabled").Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("enabled setting rolled back, count=%d", n)
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true,"generate":true}`, csrf())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("enable+generate: %d %+v", res.StatusCode, env)
}
_ = json.Unmarshal(env.Data, &got)
code, _ := got["code"].(string)
if got["enabled"] != true || len(code) != 16 {
t.Fatalf("after generate: %v", got)
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":false}`, csrf())
if res.StatusCode != 200 {
t.Fatalf("disable: %d", res.StatusCode)
}
_ = res.Body.Close()
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true,"code":"abcdefgh"}`, csrf())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("enable+code: %d %+v", res.StatusCode, env)
}
_ = json.Unmarshal(env.Data, &got)
if got["enabled"] != true || got["code"] != "abcdefgh" {
t.Fatalf("after enable+code: %v", got)
}
}
func TestMessageDetailHasNoBody(t *testing.T) { func TestMessageDetailHasNoBody(t *testing.T) {
db, srv, client := setupA3(t) db, srv, client := setupA3(t)
base := srv.URL base := srv.URL
@@ -212,6 +272,33 @@ func TestGroupsCRUD(t *testing.T) {
t.Fatalf("created=%+v", created) 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, res = doReq(t, client, http.MethodPatch, base+"/api/admin/groups/"+created.ID,
`{"name":"新名"}`, csrf()) `{"name":"新名"}`, csrf())
env = decodeEnv(t, res) env = decodeEnv(t, res)
+61 -5
View File
@@ -233,7 +233,7 @@ func (h *Handler) handleEndpointCreate(w http.ResponseWriter, r *http.Request) {
if req.DefaultDelaySeconds != nil { if req.DefaultDelaySeconds != nil {
delaySec = *req.DefaultDelaySeconds 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) h.audit(actorString(p), "endpoint_create", req.ID, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg) httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg)
return return
@@ -341,17 +341,24 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) {
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "备注过长") httpx.WriteError(w, http.StatusBadRequest, "bad_request", "备注过长")
return return
} }
if req.DefaultDelaySeconds != nil && *req.DefaultDelaySeconds < 0 { if req.DefaultDelaySeconds != nil {
if msg := validateDelaySeconds(*req.DefaultDelaySeconds, h.maxScheduleSeconds()); msg != "" {
h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "默认延迟无效") httpx.WriteError(w, http.StatusBadRequest, "bad_request", msg)
return return
} }
}
hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil
var wasEnabled bool var wasEnabled bool
var err error var err error
// 注入 Identity 时启停只走 identity,避免先写 enabled 再级联失败造成半生效。
patchEnabled := req.Enabled
if h.identity != nil {
patchEnabled = nil
}
if hasMeta || (req.Enabled != nil && h.identity == 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 err != nil {
if errors.Is(err, sql.ErrNoRows) { if errors.Is(err, sql.ErrNoRows) {
h.audit(actorString(p), "endpoint_patch", id, "not_found", ip) h.audit(actorString(p), "endpoint_patch", id, "not_found", ip)
@@ -563,6 +570,29 @@ func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Requ
return return
} }
if h.identity != nil {
err := h.identity.SelfSetTalkPassword(r.Context(), id, req.TalkPassword)
if err != nil {
if isEndpointNotFound(err) {
h.audit(actorString(p), "endpoint_talk_password", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return
}
var pe *protocol.Error
if errors.As(err, &pe) && pe.Code == protocol.CodeBadRequest {
h.audit(actorString(p), "endpoint_talk_password", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "对话密码不合法")
return
}
h.audit(actorString(p), "endpoint_talk_password", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
h.audit(actorString(p), "endpoint_talk_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"talk_password_set": req.TalkPassword != ""})
return
}
var talkHash sql.NullString var talkHash sql.NullString
if req.TalkPassword != "" { if req.TalkPassword != "" {
th, err := h.hash.Hash(r.Context(), auth.PasswordTalk, req.TalkPassword) th, err := h.hash.Hash(r.Context(), auth.PasswordTalk, req.TalkPassword)
@@ -584,6 +614,7 @@ func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Requ
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return return
} }
h.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: id})
h.audit(actorString(p), "endpoint_talk_password", id, "ok", ip) h.audit(actorString(p), "endpoint_talk_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"talk_password_set": talkHash.Valid}) httpx.WriteOK(w, map[string]any{"talk_password_set": talkHash.Valid})
} }
@@ -609,7 +640,7 @@ func (h *Handler) handleEndpointUnlock(w http.ResponseWriter, r *http.Request) {
httpx.WriteOK(w, map[string]any{}) 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) { if id != "" && !protocol.ValidEndpointID(id) {
return "编号不合法" return "编号不合法"
} }
@@ -628,12 +659,26 @@ func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec i
if !protocol.ValidTalkPassword(talkPW) { if !protocol.ValidTalkPassword(talkPW) {
return "对话密码不合法" return "对话密码不合法"
} }
if msg := validateDelaySeconds(delaySec, maxDelaySec); msg != "" {
return msg
}
return ""
}
func validateDelaySeconds(delaySec, maxDelaySec int64) string {
if delaySec < 0 { if delaySec < 0 {
return "默认延迟无效" return "默认延迟无效"
} }
if maxDelaySec > 0 && delaySec > maxDelaySec {
return "默认延迟无效"
}
return "" return ""
} }
func (h *Handler) maxScheduleSeconds() int64 {
return int64(h.cfg.Limits.MaxScheduleSeconds)
}
func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{ httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{
OK: false, OK: false,
@@ -644,3 +689,14 @@ func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
Data: map[string]any{"errors": errs}, 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},
})
}
+189 -68
View File
@@ -5,12 +5,15 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/csv" "encoding/csv"
"errors"
"io" "io"
"mime" "mime"
"mime/multipart" "mime/multipart"
"net/http" "net/http"
"runtime"
"strconv" "strconv"
"strings" "strings"
"sync"
"unicode/utf8" "unicode/utf8"
"git.asio.asia/nixevol/NixMsg/internal/auth" "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) raw, err := readImportCSV(r)
if err != nil { 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) h.audit(actorString(p), "endpoint_import", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error()) httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error())
return 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 { if len(errs) > 0 {
h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) h.audit(actorString(p), "endpoint_import", "", "bad_request", ip)
writeCSVValidationError(w, errs) 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 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) h.audit(actorString(p), "endpoint_import", "", "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return return
@@ -91,9 +113,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
} }
name := part.FormName() name := part.FormName()
if name == "file" || name == "" { if name == "file" || name == "" {
b, readErr := io.ReadAll(io.LimitReader(part, 8<<20)) b, readErr := io.ReadAll(part)
_ = part.Close() _ = part.Close()
if readErr != nil { if readErr != nil {
if httpx.IsBodyTooLarge(readErr) {
return nil, readErr
}
return nil, errBadRequest("读取文件失败") return nil, errBadRequest("读取文件失败")
} }
return b, nil return b, nil
@@ -102,9 +127,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
} }
return nil, errBadRequest("缺少 file 字段") return nil, errBadRequest("缺少 file 字段")
default: default:
// text/csv 或未标明时按原始体 // text/csv 或未标明时按原始体;大小由 ServeHTTP 的 MaxBytesReader 限制。
b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20)) b, readErr := io.ReadAll(r.Body)
if readErr != nil { if readErr != nil {
if httpx.IsBodyTooLarge(readErr) {
return nil, readErr
}
return nil, errBadRequest("读取 CSV 失败") return nil, errBadRequest("读取 CSV 失败")
} }
return b, nil return b, nil
@@ -117,45 +145,7 @@ func (e badRequestError) Error() string { return string(e) }
func errBadRequest(msg string) error { return badRequestError(msg) } func errBadRequest(msg string) error { return badRequestError(msg) }
func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError) { type importPending struct {
raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF})
reader := csv.NewReader(bytes.NewReader(raw))
reader.FieldsPerRecord = -1
reader.TrimLeadingSpace = true
records, err := reader.ReadAll()
if err != nil {
return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}}
}
if len(records) < 1 {
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
}
header := normalizeCSVHeader(records[0])
expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"}
if len(header) < len(expected) {
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
}
for i, want := range expected {
if header[i] != want {
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
}
}
dataRows := records[1:]
if len(dataRows) == 0 {
return nil, []csvLineError{{Line: 2, Reason: "没有数据行"}}
}
if len(dataRows) > maxImportRows {
return nil, []csvLineError{{Line: 1, Reason: "最多 1000 行"}}
}
errs := make([]csvLineError, 0)
prepared := make([]importPrepared, 0, len(dataRows))
seen := make(map[string]int) // id -> first line
checkIDs := make([]string, 0, len(dataRows))
type pending struct {
line int line int
id string id string
name string name string
@@ -165,11 +155,76 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
delaySec int64 delaySec int64
needGenerateID bool needGenerateID bool
needGenerateLogin bool needGenerateLogin bool
} }
pendings := make([]pending, 0, len(dataRows))
for i, cols := range dataRows { func csvErrorLine(err error, fallback int) int {
line := i + 2 // 表头为 1 var pe *csv.ParseError
if errors.As(err, &pe) {
if pe.StartLine > 0 {
return pe.StartLine
}
if pe.Line > 0 {
return pe.Line
}
}
if fallback > 0 {
return fallback
}
return 1
}
func importHashConcurrency() int {
n := runtime.NumCPU() - 1
if n < 1 {
return 1
}
return n
}
func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError, error) {
raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF})
reader := csv.NewReader(bytes.NewReader(raw))
reader.FieldsPerRecord = -1
reader.TrimLeadingSpace = true
reader.ReuseRecord = false
header, err := reader.Read()
if err != nil {
return nil, []csvLineError{{Line: csvErrorLine(err, 1), Reason: "CSV 解析失败"}}, nil
}
headerLine, _ := reader.FieldPos(0)
if headerLine < 1 {
headerLine = 1
}
norm := normalizeCSVHeader(header)
expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"}
if len(norm) < len(expected) {
return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil
}
for i, want := range expected {
if norm[i] != want {
return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil
}
}
errs := make([]csvLineError, 0)
seen := make(map[string]int)
checkIDs := make([]string, 0)
pendings := make([]importPending, 0)
for {
cols, readErr := reader.Read()
if errors.Is(readErr, io.EOF) {
break
}
if readErr != nil {
errs = append(errs, csvLineError{Line: csvErrorLine(readErr, 1), Reason: "CSV 解析失败"})
continue
}
line, _ := reader.FieldPos(0)
if line < 1 {
line = headerLine + 1
}
for len(cols) < 6 { for len(cols) < 6 {
cols = append(cols, "") cols = append(cols, "")
} }
@@ -189,7 +244,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
} }
delaySec = n 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}) errs = append(errs, csvLineError{Line: line, Reason: msg})
continue continue
} }
@@ -198,7 +253,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
continue continue
} }
p := pending{ p := importPending{
line: line, line: line,
id: id, id: id,
name: name, name: name,
@@ -221,12 +276,18 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
} }
if len(errs) > 0 { 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) existing, err := h.existingEndpointIDs(ctx, checkIDs)
if err != nil { if err != nil {
return nil, []csvLineError{{Line: 1, Reason: "校验失败"}} return nil, nil, err
} }
for _, p := range pendings { for _, p := range pendings {
if p.needGenerateID { if p.needGenerateID {
@@ -237,16 +298,19 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
} }
} }
if len(errs) > 0 { if len(errs) > 0 {
return nil, errs return nil, errs, nil
} }
for _, p := range pendings { hashed := make([]importPending, len(pendings))
useID := p.id copy(hashed, pendings)
for i := range hashed {
p := &hashed[i]
if p.needGenerateID { if p.needGenerateID {
useID := ""
for attempt := 0; attempt < 16; attempt++ { for attempt := 0; attempt < 16; attempt++ {
genID, genErr := generateEndpointID() genID, genErr := generateEndpointID()
if genErr != nil { if genErr != nil {
return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} return nil, nil, genErr
} }
if _, clash := seen[genID]; clash { if _, clash := seen[genID]; clash {
continue continue
@@ -259,33 +323,62 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
break break
} }
if useID == "" { 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 { if p.needGenerateLogin {
pw, genErr := generateLoginPassword() pw, genErr := generateLoginPassword()
if genErr != nil { 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) }
prepared, hashErr := h.hashImportPendings(ctx, hashed)
if hashErr != nil { if hashErr != nil {
return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}} 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 var talkHash sql.NullString
if p.talkPW != "" { if p.talkPW != "" {
th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW) th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW)
if thErr != nil { if thErr != nil {
return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}} errCh[i] = thErr
return
} }
talkHash = sql.NullString{String: th, Valid: true} talkHash = sql.NullString{String: th, Valid: true}
} }
prepared = append(prepared, importPrepared{ out[i] = importPrepared{
Insert: endpointInsert{ Insert: endpointInsert{
ID: useID, ID: p.id,
Name: p.name, Name: p.name,
Remark: p.remark, Remark: p.remark,
Source: sourceAdmin, Source: sourceAdmin,
@@ -293,12 +386,40 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
TalkHash: talkHash, TalkHash: talkHash,
DefaultDelayMs: p.delaySec * 1000, DefaultDelayMs: p.delaySec * 1000,
}, },
PlainLogin: loginPW, PlainLogin: p.loginPW,
Name: p.name, Name: p.name,
Line: p.line, Line: p.line,
})
} }
return prepared, nil }()
}
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 { func normalizeCSVHeader(cols []string) []string {
+3
View File
@@ -370,6 +370,9 @@ func (h *Handler) deleteEndpointBasic(ctx context.Context, id string) (found boo
found = n > 0 found = n > 0
return nil return nil
}) })
if err == nil && found {
h.locks.ClearAllForEndpoint(id)
}
return found, err return found, err
} }
+10
View File
@@ -269,6 +269,13 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
t.Fatal("expected unlocked") t.Fatal("expected unlocked")
} }
for i := 0; i < 50; i++ {
locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"})
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"}); !locked {
t.Fatal("expected talk target lock before talk-password")
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password", res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password",
`{"talk_password":"talk"}`, `{"talk_password":"talk"}`,
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"}) map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
@@ -283,6 +290,9 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
if !talkSet.TalkPasswordSet { if !talkSet.TalkPasswordSet {
t.Fatal("want talk_password_set true") t.Fatal("want talk_password_set true")
} }
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"}); locked {
t.Fatal("talk-password should clear LockTalkTarget")
}
var ver int var ver int
if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil { if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil {
+1
View File
@@ -209,6 +209,7 @@ LIMIT ? OFFSET ?`, id, limit, offset)
"name": name, "name": name,
"owner_id": owner, "owner_id": owner,
"created_at_ms": created, "created_at_ms": created,
"member_total": memberTotal,
"members": members, "members": members,
"next_cursor": next, "next_cursor": next,
}) })
+53
View File
@@ -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)
}
}
+146
View File
@@ -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)
}
}
+276
View File
@@ -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)
}
}
+61
View File
@@ -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)
}
}
+75
View File
@@ -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)
}
}
+35 -1
View File
@@ -1,6 +1,7 @@
package admin package admin
import ( import (
"errors"
"log/slog" "log/slog"
"net" "net"
"net/http" "net/http"
@@ -11,6 +12,7 @@ import (
"git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/app/identity"
"git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/httpx"
"git.asio.asia/nixevol/NixMsg/internal/store" "git.asio.asia/nixevol/NixMsg/internal/store"
) )
@@ -23,6 +25,13 @@ const (
defaultSessionTTL = 12 * time.Hour defaultSessionTTL = 12 * time.Hour
minPasswordLen = 12 minPasswordLen = 12
lastUsedMinGap = time.Minute 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 的依赖。 // Deps 是管理 Handler 的依赖。
@@ -129,11 +138,36 @@ func New(d Deps) *Handler {
return h return h
} }
// ServeHTTP 实现 http.Handler。 // ServeHTTP 实现 http.Handler。按路由限制请求体大小与读截止时间。
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { 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) 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() { func (h *Handler) routes() {
// 公开 // 公开
h.mux.HandleFunc("POST /api/admin/login", h.handleLogin) h.mux.HandleFunc("POST /api/admin/login", h.handleLogin)
+25 -5
View File
@@ -4,6 +4,7 @@ import (
"database/sql" "database/sql"
"errors" "errors"
"net/http" "net/http"
"unicode/utf8"
"git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/httpx" "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"` Password string `json:"password"`
} }
if err := httpx.DecodeJSON(r, &req); err != nil { if err := httpx.DecodeJSON(r, &req); err != nil {
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") writeDecodeError(w, err)
return return
} }
if req.Username != adminUsername { if req.Username != adminUsername {
@@ -111,16 +112,23 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
p, _ := principalFrom(r.Context()) p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted) 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 { var req struct {
OldPassword string `json:"old_password"` OldPassword string `json:"old_password"`
NewPassword string `json:"new_password"` NewPassword string `json:"new_password"`
} }
if err := httpx.DecodeJSON(r, &req); err != nil { if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "password_change", "", "bad_request", ip) h.audit(actorString(p), "password_change", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") writeDecodeError(w, err)
return return
} }
if len(req.NewPassword) < minPasswordLen { if utf8.RuneCountInString(req.NewPassword) < minPasswordLen {
h.audit(actorString(p), "password_change", "", "bad_request", ip) h.audit(actorString(p), "password_change", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "新密码至少 12 位") httpx.WriteError(w, http.StatusBadRequest, "bad_request", "新密码至少 12 位")
return 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) ok, err := h.hash.Verify(r.Context(), auth.PasswordAdmin, req.OldPassword, phc)
if err != nil || !ok { 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) h.audit(actorString(p), "password_change", "", "unauthorized", ip)
httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "旧密码错误") httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "旧密码错误")
return return
} }
h.locks.Clear(auth.LockKey{Kind: auth.LockAdminIP, IP: ip})
newPHC, err := h.hash.Hash(r.Context(), auth.PasswordAdmin, req.NewPassword) newPHC, err := h.hash.Hash(r.Context(), auth.PasswordAdmin, req.NewPassword)
if err != nil { if err != nil {
h.audit(actorString(p), "password_change", "", "error", ip) 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", "内部错误") httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return return
} }
// 保留当前会话,作废其它会话 // 保留当前会话,作废其它会话;失败则返回 500,避免其它会话继续有效。
if p.Session != "" { 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) h.audit(actorString(p), "password_change", "", "ok", ip)
httpx.WriteOK(w, map[string]any{}) httpx.WriteOK(w, map[string]any{})
+27 -3
View File
@@ -89,14 +89,38 @@ func (l *MemoryLoginLocks) Fail(key auth.LockKey) (bool, time.Duration) {
return false, 0 return false, 0
} }
// ClearEndpoint 实现 auth.LoginLocks。 // ClearEndpoint 实现 auth.LoginLocks:只清登录锁定。
func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) { func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
for k := range l.entries { for k := range l.entries {
// kind|endpoint|peer|ip
parts := splitLockKey(k) parts := splitLockKey(k)
if len(parts) >= 2 && parts[1] == endpointID { if len(parts) != 4 {
continue
}
kind, ep := auth.LockKind(parts[0]), parts[1]
if ep != endpointID {
continue
}
if kind == auth.LockLoginEndpointIP || kind == auth.LockLoginEndpoint {
delete(l.entries, k)
}
}
}
// ClearAllForEndpoint 实现 auth.LoginLocks。
func (l *MemoryLoginLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
if endpointID == "" {
return
}
for k := range l.entries {
parts := splitLockKey(k)
if len(parts) != 4 {
continue
}
if parts[1] == endpointID || parts[2] == endpointID {
delete(l.entries, k) delete(l.entries, k)
} }
} }
+13 -3
View File
@@ -77,9 +77,7 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request)
); e != nil { ); e != nil {
return e return e
} }
return nil } else if req.Code != nil {
}
if req.Code != nil {
code := *req.Code code := *req.Code
n := utf8.RuneCountInString(code) n := utf8.RuneCountInString(code)
if n < minRegistrationCodeLen || n > maxRegistrationCodeLen { if n < minRegistrationCodeLen || n > maxRegistrationCodeLen {
@@ -93,6 +91,18 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request)
return e return e
} }
} }
var enabledVal, codeVal sql.NullString
if e := tx.QueryRow(`SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled).Scan(&enabledVal); e != nil && !errors.Is(e, sql.ErrNoRows) {
return e
}
if e := tx.QueryRow(`SELECT value FROM settings WHERE key = ?`, settingRegistrationCode).Scan(&codeVal); e != nil && !errors.Is(e, sql.ErrNoRows) {
return e
}
n := utf8.RuneCountInString(codeVal.String)
if registrationTruthy(enabledVal.String) && (n < minRegistrationCodeLen || n > maxRegistrationCodeLen) {
return errBadRequest("开启自助注册前须先设置 8–64 字符的安全码")
}
return nil return nil
}) })
if err != nil { if err != nil {
+3
View File
@@ -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 { func (h *Handler) createSession(ctx context.Context, hashHex string, ttl time.Duration) error {
now := time.Now() now := time.Now()
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error { 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( _, err := tx.Exec(
`INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES(?, ?, ?)`, `INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES(?, ?, ?)`,
hashHex, now.UnixMilli(), now.Add(ttl).UnixMilli(), hashHex, now.UnixMilli(), now.Add(ttl).UnixMilli(),
+3
View File
@@ -126,6 +126,9 @@ WHERE id = ?`, endpointID); e != nil {
if err != nil { if err != nil {
return err return err
} }
if hardDelete {
a.locks.ClearAllForEndpoint(endpointID)
}
a.publishRevokes(ctx, revokes) a.publishRevokes(ctx, revokes)
a.publishGroupEvents(ctx, notifies) a.publishGroupEvents(ctx, notifies)
+92
View File
@@ -501,6 +501,98 @@ func TestAdminDisableDeleteHTTP(t *testing.T) {
} }
} }
func TestU03AdminTalkPasswordHTTP(t *testing.T) {
t.Parallel()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
hash := auth.NewStubHashPool()
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, "adminpassword1"); seedErr != nil {
t.Fatal(seedErr)
}
locks := auth.NewLoginLocks()
locks.SetClock(func() time.Time { return fixed })
idApp := identity.New(identity.Config{
DB: db, Hash: hash, Locks: locks, Sessions: auth.NewSessionTokens(),
MaxScheduleSeconds: 86400, Now: func() time.Time { return fixed },
})
h := admin.New(admin.Deps{
DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(),
Locks: locks, Identity: idApp,
})
srv := httptest.NewServer(h)
t.Cleanup(srv.Close)
jar, _ := cookiejar.New(nil)
client := &http.Client{Jar: jar}
loginRes, err := client.Post(srv.URL+"/api/admin/login", "application/json",
strings.NewReader(`{"username":"admin","password":"adminpassword1"}`))
if err != nil {
t.Fatal(err)
}
_ = loginRes.Body.Close()
if loginRes.StatusCode != 200 {
t.Fatalf("login %d", loginRes.StatusCode)
}
req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints",
strings.NewReader(`{"id":"alice","name":"alice","login_password":"password12"}`))
req.Header.Set("Content-Type", "application/json")
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 != 200 {
t.Fatalf("create: %d %s", res.StatusCode, raw)
}
for i := 0; i < 50; i++ {
locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"})
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); !locked {
t.Fatal("expected target lock")
}
req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/alice/talk-password",
strings.NewReader(`{"talk_password":"secret"}`))
req.Header.Set("Content-Type", "application/json")
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 != 200 {
t.Fatalf("talk-password: %d %s", res.StatusCode, raw)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); locked {
t.Fatal("identity SelfSetTalkPassword should clear LockTalkTarget")
}
req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/missing/talk-password",
strings.NewReader(`{"talk_password":"secret"}`))
req.Header.Set("Content-Type", "application/json")
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.StatusNotFound {
t.Fatalf("missing endpoint want 404 got %d %s", res.StatusCode, raw)
}
}
type lifecycleKick struct { type lifecycleKick struct {
Calls []string Calls []string
} }
+11
View File
@@ -25,6 +25,9 @@ const (
settingRegistrationEnabled = "registration_enabled" settingRegistrationEnabled = "registration_enabled"
settingRegistrationCode = "registration_code" settingRegistrationCode = "registration_code"
minRegistrationCodeLen = 8
maxRegistrationCodeLen = 64
sourceSelf = "self" sourceSelf = "self"
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789" idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
@@ -144,6 +147,11 @@ func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRe
if !enabled { if !enabled {
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed") return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
} }
// 存储码不是 8–64 字符时视为关闭(含开启+空码的旧库)。放在锁定检查之前,不计入锁定。
if n := utf8.RuneCountInString(storedCode); n < minRegistrationCodeLen || n > maxRegistrationCodeLen {
h.cfg.Logger.Warn("register", "result", protocol.CodeRegistrationClosed, "reason", "unusable_code", "ip", ip)
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
}
lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip} lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip}
if locked, _ := h.cfg.Locks.Check(lockKey); locked { if locked, _ := h.cfg.Locks.Check(lockKey); locked {
@@ -163,6 +171,9 @@ func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRe
} }
return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg) return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg)
} }
if strings.EqualFold(req.ID, "inline") {
return RegisterResult{}, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "id reserved")
}
id := req.ID id := req.ID
loginPassword := req.LoginPassword loginPassword := req.LoginPassword
+77
View File
@@ -93,6 +93,7 @@ func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) {
} }
func (l *registerIPLocker) ClearEndpoint(string) {} func (l *registerIPLocker) ClearEndpoint(string) {}
func (l *registerIPLocker) ClearAllForEndpoint(string) {}
func (l *registerIPLocker) Clear(key auth.LockKey) { func (l *registerIPLocker) Clear(key auth.LockKey) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -141,6 +142,33 @@ func openTestEnv(t *testing.T) *testEnv {
return env return env
} }
func (e *testEnv) setEnabledOnly(t *testing.T, enabled bool) {
t.Helper()
en := "0"
if enabled {
en = "1"
}
now := time.Now().UnixMilli()
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
if _, err := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
settingRegistrationEnabled, en, now); err != nil {
return err
}
_, err := tx.Exec(`DELETE FROM settings WHERE key = ?`, settingRegistrationCode)
return err
})
if err != nil {
t.Fatal(err)
}
}
func (l *registerIPLocker) failCount(ip string) int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.fails[ip])
}
func (e *testEnv) setRegistration(t *testing.T, enabled bool, code string) { func (e *testEnv) setRegistration(t *testing.T, enabled bool, code string) {
t.Helper() t.Helper()
en := "0" en := "0"
@@ -233,6 +261,43 @@ func TestRegisterF23_ClosedFails(t *testing.T) {
} }
} }
func TestRegister_EnabledWithoutUsableCodeIsClosed(t *testing.T) {
cases := []struct {
name string
prep func(*testEnv, *testing.T)
}{
{"no_code_row", func(env *testEnv, t *testing.T) { env.setEnabledOnly(t, true) }},
{"empty_code", func(env *testEnv, t *testing.T) { env.setRegistration(t, true, "") }},
{"short_code", func(env *testEnv, t *testing.T) { env.setRegistration(t, true, "short") }},
}
bodies := []string{
`{"registration_code":"","id":"ep_empty","login_password":"password1"}`,
`{"registration_code":"anything1","id":"ep_any","login_password":"password1"}`,
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
env := openTestEnv(t)
tc.prep(env, t)
for _, body := range bodies {
code, resp, _ := env.doRegister(t, body)
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationClosed {
t.Fatalf("body=%s status=%d resp=%+v", body, code, resp)
}
}
if env.locks.failCount(env.fixedIP) != 0 {
t.Fatalf("unusable code must not count as lock fail: %d", env.locks.failCount(env.fixedIP))
}
logged := env.logBuf.String()
if strings.Contains(logged, "anything1") || strings.Contains(logged, "short") {
t.Fatalf("log leaked code: %s", logged)
}
if !strings.Contains(logged, "unusable_code") {
t.Fatalf("missing warn: %s", logged)
}
})
}
}
func TestRegisterF23_WrongCodeFails_RightCodeOK(t *testing.T) { func TestRegisterF23_WrongCodeFails_RightCodeOK(t *testing.T) {
env := openTestEnv(t) env := openTestEnv(t)
env.setRegistration(t, true, "good-code-01") env.setRegistration(t, true, "good-code-01")
@@ -475,6 +540,18 @@ func TestRegister_LogOmitsSecrets(t *testing.T) {
} }
} }
func TestRegister_ReservedInlineID(t *testing.T) {
env := openTestEnv(t)
env.setRegistration(t, true, "inline-code1")
code, resp, _ := env.doRegister(t, `{"registration_code":"inline-code1","id":"inline","login_password":"password1"}`)
if code != http.StatusBadRequest || resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
t.Fatalf("status=%d resp=%+v", code, resp)
}
if _, _, ok := env.getEndpoint(t, "inline"); ok {
t.Fatal("inline must not be created")
}
}
func TestRegister_NSTPasswordRejected(t *testing.T) { func TestRegister_NSTPasswordRejected(t *testing.T) {
env := openTestEnv(t) env := openTestEnv(t)
env.setRegistration(t, true, "nst-code-01") env.setRegistration(t, true, "nst-code-01")
+102
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"fmt"
"path/filepath" "path/filepath"
"testing" "testing"
"time" "time"
@@ -297,6 +298,107 @@ func TestSelfLoginPasswordAndLogout(t *testing.T) {
} }
} }
func TestU03TalkPairLockIgnoresIPAndEmptyPassword(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
for i := 0; i < 9; i++ {
err := app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("10.0.0.%d", i+1))
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("fail %d: %v", i, err)
}
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "9.9.9.9"); err != nil {
t.Fatalf("9th+correct from other IP should succeed: %v", err)
}
for i := 0; i < 5; i++ {
if err := app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("1.1.1.%d", i+1)); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("unlock fail %d: %v", i, err)
}
}
for i := 0; i < 5; i++ {
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "wrong", fmt.Sprintf("2.2.2.%d", i+1)); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("join fail %d: %v", i, err)
}
}
pair := auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "alice", PeerID: "bob"}
if locked, _ := locks.Check(pair); !locked {
t.Fatal("pair should lock after 10 wrong attempts across IPs and paths")
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "8.8.8.8"); protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited unlock got %v", err)
}
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "secret", "7.7.7.7"); protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited join got %v", err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "", "6.6.6.6"); protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("empty while locked want required got %v", err)
}
}
func TestU03AdminChangeAndDeleteClearTalkLocks(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
for i := 0; i < 50; i++ {
id := fmt.Sprintf("u%02d", i)
insertEP(t, db, id, "password1")
_ = app.UnlockTalk(ctx, id, "bob", "wrong", "2.2.2.2")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); !locked {
t.Fatal("expected talk target lock")
}
if err := app.SelfSetTalkPassword(ctx, "bob", "secret2"); err != nil {
t.Fatal(err)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); locked {
t.Fatal("admin/self change should clear LockTalkTarget")
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret2", "3.3.3.3"); err != nil {
t.Fatalf("unlock after change: %v", err)
}
for i := 0; i < 10; i++ {
_ = app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("4.4.4.%d", i+1))
}
locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "bob"})
locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "bob", PeerID: "alice"})
if err := app.Delete(ctx, "bob"); err != nil {
t.Fatal(err)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "alice", PeerID: "bob"}); locked {
t.Fatal("delete should clear pair where bob is peer")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "bob", PeerID: "alice"}); locked {
t.Fatal("delete should clear pair where bob is sender")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); locked {
t.Fatal("delete should clear LockTalkTarget")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "bob"}); locked {
t.Fatal("delete should clear login lock")
}
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "5.5.5.5"); err != nil {
t.Fatalf("reopened bob must not inherit talk lock: %v", err)
}
}
func TestUnlockSelfAndNoPassword(t *testing.T) { func TestUnlockSelfAndNoPassword(t *testing.T) {
t.Parallel() t.Parallel()
app, db, _ := openIdentity(t) app, db, _ := openIdentity(t)
+2 -1
View File
@@ -50,7 +50,8 @@ type Service interface {
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
SelfLogout(ctx context.Context, endpointID string) error SelfLogout(ctx context.Context, endpointID string) error
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。 // UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock)。
// remoteIP 保留给接线方;对话密码锁键为 {LockTalkPair, 发送方, 对方},不含 IP。
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。 // HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error)
+47 -18
View File
@@ -32,15 +32,13 @@ SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&tal
return nil return nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked { // 没带密码不计锁定、也不因已锁返回 rate_limited(DEVELOPMENT 第 5 节:带密码的发送/进群才限流)。
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if talkPassword == "" { if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required") return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
} }
if a.talkRateLimited(senderID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil { if a.hash == nil {
return errors.New("identity: hash pool required") return errors.New("identity: hash pool required")
} }
@@ -49,11 +47,11 @@ SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&tal
return err return err
} }
if !ok { if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) a.talkFail(senderID, targetID)
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
} }
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) a.talkClearPair(senderID, targetID)
_ = remoteIP
nowMs := a.now().UnixMilli() nowMs := a.now().UnixMilli()
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
@@ -116,15 +114,12 @@ SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).
return nil return nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if talkPassword == "" { if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required") return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
} }
if a.talkRateLimited(actorID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil { if a.hash == nil {
return errors.New("identity: hash pool required") return errors.New("identity: hash pool required")
} }
@@ -133,16 +128,50 @@ SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).
return err return err
} }
if !ok { if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) a.talkFail(actorID, targetID)
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
} }
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) a.talkClearPair(actorID, targetID)
// 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。 // 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。
_ = talkVer _ = talkVer
_ = remoteIP
return nil return nil
} }
func talkPairKey(senderID, targetID string) auth.LockKey {
return auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID}
}
func talkTargetKey(targetID string) auth.LockKey {
return auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}
}
func (a *App) talkRateLimited(senderID, targetID string) bool {
if a.locks == nil {
return false
}
if locked, _ := a.locks.Check(talkPairKey(senderID, targetID)); locked {
return true
}
locked, _ := a.locks.Check(talkTargetKey(targetID))
return locked
}
func (a *App) talkFail(senderID, targetID string) {
if a.locks == nil {
return
}
a.locks.Fail(talkPairKey(senderID, targetID))
a.locks.Fail(talkTargetKey(targetID))
}
func (a *App) talkClearPair(senderID, targetID string) {
if a.locks == nil {
return
}
a.locks.Clear(talkPairKey(senderID, targetID))
}
func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error { func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
_, err := tx.Exec(` _, err := tx.Exec(`
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at) INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at)
+8 -8
View File
@@ -118,12 +118,12 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
} }
passwordVerified := false passwordVerified := false
if needPassword { if needPassword {
if locked, _ := a.talkLocked(senderID, req.To.ID, conn.RemoteIP); locked {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
}
if req.TalkPassword == "" { if req.TalkPassword == "" {
return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required") return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required")
} }
if locked, _ := a.talkLocked(senderID, req.To.ID); locked {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil { if a.hash == nil {
return SubmitResult{}, fmt.Errorf("message: hash pool required") return SubmitResult{}, fmt.Errorf("message: hash pool required")
} }
@@ -132,7 +132,7 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
return SubmitResult{}, vErr return SubmitResult{}, vErr
} }
if !ok { if !ok {
a.talkFail(senderID, req.To.ID, conn.RemoteIP) a.talkFail(senderID, req.To.ID)
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
} }
passwordVerified = true passwordVerified = true
@@ -569,11 +569,11 @@ func bytesEqual(a, b []byte) bool {
return v == 0 return v == 0
} }
func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) { func (a *App) talkLocked(senderID, targetID string) (bool, error) {
if a.locks == nil { if a.locks == nil {
return false, nil return false, nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip}); locked { if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID}); locked {
return true, nil return true, nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
@@ -582,11 +582,11 @@ func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
return false, nil return false, nil
} }
func (a *App) talkFail(senderID, targetID, ip string) { func (a *App) talkFail(senderID, targetID string) {
if a.locks == nil { if a.locks == nil {
return return
} }
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip}) a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID})
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}) a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
} }
+42
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"fmt"
"math" "math"
"path/filepath" "path/filepath"
"testing" "testing"
@@ -519,3 +520,44 @@ WHERE m.id='late-1' AND d.endpoint_id='bob'`).Scan(&bobN); err != nil {
} }
}) })
} }
func TestU03SubmitTalkLockNoIP(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
locks := auth.NewLoginLocks()
locks.SetClock(func() time.Time { return fixed })
app := New(db, lim, auth.NewStubHashPool(),
WithNow(func() time.Time { return fixed }),
WithLocks(locks),
)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "secret", 1, 0)
ctx := context.Background()
for i := 0; i < 10; i++ {
bad := baseSend(fmt.Sprintf("w%d", i), "bob")
bad.TalkPassword = "wrong"
_, err := app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: fmt.Sprintf("10.0.0.%d", i+1)}, bad)
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("i=%d got %v", i, err)
}
}
empty := baseSend("empty", "bob")
_, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "8.8.8.8"}, empty)
if protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("empty while locked want required got %v", err)
}
okReq := baseSend("ok1", "bob")
okReq.TalkPassword = "secret"
_, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "9.9.9.9"}, okReq)
if protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("correct password while locked want rate_limited got %v", err)
}
}
+5 -3
View File
@@ -122,13 +122,15 @@ WHERE id > ?
ORDER BY id ASC ORDER BY id ASC
LIMIT ?`, cursor, limit+1) LIMIT ?`, cursor, limit+1)
} else { } else {
like := "%" + strings.ToLower(query) + "%" q := strings.ToLower(query)
prefix := strings.ToLower(query) + "%" esc := escapeLikePattern(q)
like := "%" + esc + "%"
prefix := esc + "%"
rows, err = a.db.Read.QueryContext(ctx, ` rows, err = a.db.Read.QueryContext(ctx, `
SELECT id, name, online_since, offline_since, talk_hash SELECT id, name, online_since, offline_since, talk_hash
FROM endpoints FROM endpoints
WHERE id > ? WHERE id > ?
AND (lower(id) LIKE ? OR lower(name) LIKE ?) AND (lower(id) LIKE ? ESCAPE '\' OR lower(name) LIKE ? ESCAPE '\')
ORDER BY id ASC ORDER BY id ASC
LIMIT ?`, cursor, prefix, like, limit+1) LIMIT ?`, cursor, prefix, like, limit+1)
} }
+16
View File
@@ -0,0 +1,16 @@
package presence
import "strings"
func escapeLikePattern(s string) string {
var b strings.Builder
b.Grow(len(s) + 4)
for _, r := range s {
switch r {
case '\\', '%', '_':
b.WriteByte('\\')
}
b.WriteRune(r)
}
return b.String()
}
+29
View File
@@ -131,6 +131,35 @@ func TestF03PresenceAndDirectory(t *testing.T) {
} }
} }
func TestDirectoryLikeEscapesUnderscore(t *testing.T) {
t.Parallel()
app, db, _ := openPresence(t)
ctx := context.Background()
insertEP(t, db, "e_ab1", "underscore")
insertEP(t, db, "exab2", "wildcard")
insertEP(t, db, "pct", "has%percent")
q, _, err := app.Directory(ctx, &protocol.DirectoryList{
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "u1", Query: "e_ab", Limit: 10,
})
if err != nil {
t.Fatal(err)
}
if len(q) != 1 || q[0].ID != "e_ab1" {
t.Fatalf("want only e_ab1, got %+v", q)
}
q, _, err = app.Directory(ctx, &protocol.DirectoryList{
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "u2", Query: "%", Limit: 10,
})
if err != nil {
t.Fatal(err)
}
if len(q) != 1 || q[0].ID != "pct" {
t.Fatalf("literal %% should not match all, got %+v", q)
}
}
func TestF04PresenceWatch(t *testing.T) { func TestF04PresenceWatch(t *testing.T) {
t.Parallel() t.Parallel()
app, db, down := openPresence(t) app, db, down := openPresence(t)
+3 -1
View File
@@ -64,7 +64,7 @@ const (
LockLoginEndpointIP LockKind = "login_endpoint_ip" LockLoginEndpointIP LockKind = "login_endpoint_ip"
// LockLoginEndpoint:编号总数,1 小时内 50 次错 → 暂停该编号密码登录 1 小时。 // LockLoginEndpoint:编号总数,1 小时内 50 次错 → 暂停该编号密码登录 1 小时。
LockLoginEndpoint LockKind = "login_endpoint" LockLoginEndpoint LockKind = "login_endpoint"
// LockTalkPair:发送方 + 对方对话密码。 // LockTalkPair:发送方 + 对方对话密码(不含 IP)。
LockTalkPair LockKind = "talk_pair" LockTalkPair LockKind = "talk_pair"
// LockTalkTarget:对方对话密码总数。 // LockTalkTarget:对方对话密码总数。
LockTalkTarget LockKind = "talk_target" LockTalkTarget LockKind = "talk_target"
@@ -90,6 +90,8 @@ type LoginLocks interface {
Fail(key LockKey) (locked bool, retryAfter time.Duration) Fail(key LockKey) (locked bool, retryAfter time.Duration)
// ClearEndpoint 清除某端编号相关的登录锁定(两种都清),对应管理 unlock。 // ClearEndpoint 清除某端编号相关的登录锁定(两种都清),对应管理 unlock。
ClearEndpoint(endpointID string) ClearEndpoint(endpointID string)
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定(删除端后防编号复用继承)。
ClearAllForEndpoint(endpointID string)
// Clear 清除精确键。 // Clear 清除精确键。
Clear(key LockKey) Clear(key LockKey)
} }
+51
View File
@@ -138,3 +138,54 @@ func TestLoginLocksClearEndpoint(t *testing.T) {
t.Fatal("e1 total lock should clear") t.Fatal("e1 total lock should clear")
} }
} }
func TestLoginLocksClearEndpointKeepsTalk(t *testing.T) {
locks := NewLoginLocks()
talk := LockKey{Kind: LockTalkPair, EndpointID: "e1", PeerID: "e2"}
target := LockKey{Kind: LockTalkTarget, EndpointID: "e1"}
login := LockKey{Kind: LockLoginEndpoint, EndpointID: "e1"}
for i := 0; i < 10; i++ {
locks.Fail(talk)
locks.Fail(login)
}
for i := 0; i < 50; i++ {
locks.Fail(target)
}
locks.ClearEndpoint("e1")
if locked, _ := locks.Check(login); locked {
t.Fatal("login lock should clear")
}
if locked, _ := locks.Check(talk); !locked {
t.Fatal("talk pair lock must survive ClearEndpoint")
}
if locked, _ := locks.Check(target); !locked {
t.Fatal("talk target lock must survive ClearEndpoint")
}
}
func TestLoginLocksClearAllForEndpoint(t *testing.T) {
locks := NewLoginLocks()
asSender := LockKey{Kind: LockTalkPair, EndpointID: "gone", PeerID: "peer"}
asPeer := LockKey{Kind: LockTalkPair, EndpointID: "other", PeerID: "gone"}
target := LockKey{Kind: LockTalkTarget, EndpointID: "gone"}
login := LockKey{Kind: LockLoginEndpointIP, EndpointID: "gone", IP: "1.1.1.1"}
keep := LockKey{Kind: LockTalkPair, EndpointID: "keep", PeerID: "peer"}
for i := 0; i < 10; i++ {
locks.Fail(asSender)
locks.Fail(asPeer)
locks.Fail(login)
locks.Fail(keep)
}
for i := 0; i < 50; i++ {
locks.Fail(target)
}
locks.ClearAllForEndpoint("gone")
for _, k := range []LockKey{asSender, asPeer, target, login} {
if locked, _ := locks.Check(k); locked {
t.Fatalf("expected %s cleared", lockMapKey(k))
}
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("unrelated pair should remain")
}
}
+134 -3
View File
@@ -21,9 +21,16 @@ func policyFor(kind LockKind) lockPolicy {
} }
} }
const (
defaultMaxLockEntries = 65536
lockSweepEveryOps = 1024
lockSweepInterval = time.Minute
)
type lockEntry struct { type lockEntry struct {
fails []time.Time fails []time.Time
lockedUntil time.Time lockedUntil time.Time
lastFail time.Time
} }
// MemoryLocks 是内存锁定计数器(重启清零)。 // MemoryLocks 是内存锁定计数器(重启清零)。
@@ -31,6 +38,9 @@ type MemoryLocks struct {
mu sync.Mutex mu sync.Mutex
entries map[string]*lockEntry entries map[string]*lockEntry
now func() time.Time now func() time.Time
ops int
lastSweep time.Time
maxEntries int
} }
// NewLoginLocks 创建默认锁定计数器。 // NewLoginLocks 创建默认锁定计数器。
@@ -38,6 +48,7 @@ func NewLoginLocks() *MemoryLocks {
return &MemoryLocks{ return &MemoryLocks{
entries: make(map[string]*lockEntry), entries: make(map[string]*lockEntry),
now: time.Now, now: time.Now,
maxEntries: defaultMaxLockEntries,
} }
} }
@@ -56,23 +67,121 @@ func lockMapKey(key LockKey) string {
return string(key.Kind) + "|" + key.EndpointID + "|" + key.PeerID + "|" + key.IP return string(key.Kind) + "|" + key.EndpointID + "|" + key.PeerID + "|" + key.IP
} }
func (l *MemoryLocks) maxCap() int {
if l.maxEntries <= 0 {
return defaultMaxLockEntries
}
return l.maxEntries
}
func (e *lockEntry) stale(now time.Time, window time.Duration) bool {
if e == nil {
return true
}
if e.lockedUntil.After(now) {
return false
}
if e.lastFail.IsZero() {
return true
}
return !e.lastFail.After(now.Add(-window))
}
func (l *MemoryLocks) maybeSweepLocked(now time.Time) {
max := l.maxCap()
if l.lastSweep.IsZero() {
l.lastSweep = now
}
l.ops++
if l.ops < lockSweepEveryOps && now.Sub(l.lastSweep) < lockSweepInterval && len(l.entries) <= max {
return
}
l.ops = 0
l.lastSweep = now
l.sweepExpiredLocked(now)
l.enforceCapLocked(now)
}
func (l *MemoryLocks) sweepExpiredLocked(now time.Time) {
for k, e := range l.entries {
window := policyFor("").window
parts := splitLockKey(k)
if len(parts) == 4 {
window = policyFor(LockKind(parts[0])).window
}
if e.stale(now, window) {
delete(l.entries, k)
}
}
}
func (l *MemoryLocks) enforceCapLocked(now time.Time) {
max := l.maxCap()
for len(l.entries) > max {
var (
victim string
victimLocked bool
victimTime time.Time
found bool
)
for k, e := range l.entries {
locked := e.lockedUntil.After(now)
t := e.lastFail
if locked {
t = e.lockedUntil
}
better := !found
if found {
if victimLocked && !locked {
better = true
} else if victimLocked == locked && (t.Before(victimTime) || (t.Equal(victimTime) && k < victim)) {
better = true
}
}
if better {
found = true
victim, victimLocked, victimTime = k, locked, t
}
}
if !found {
return
}
delete(l.entries, victim)
}
}
func (l *MemoryLocks) entryCount() int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.entries)
}
// Check 若当前已锁定返回 locked=true 与剩余时间。 // Check 若当前已锁定返回 locked=true 与剩余时间。
func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) { func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
now := l.now() now := l.now()
e := l.entries[lockMapKey(key)] k := lockMapKey(key)
e := l.entries[k]
if e == nil { if e == nil {
l.maybeSweepLocked(now)
return false, 0 return false, 0
} }
if e.lockedUntil.After(now) { if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now) return true, e.lockedUntil.Sub(now)
} }
// 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。 if e.stale(now, policyFor(key.Kind).window) {
if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) { delete(l.entries, k)
l.maybeSweepLocked(now)
return false, 0
}
// 到期自动解除:清空锁定与窗口内失败(保留仍在窗口内的失败计数)。
if !e.lockedUntil.IsZero() {
e.lockedUntil = time.Time{} e.lockedUntil = time.Time{}
e.fails = nil e.fails = nil
} }
l.maybeSweepLocked(now)
return false, 0 return false, 0
} }
@@ -88,6 +197,7 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
l.entries[k] = e l.entries[k] = e
} }
if e.lockedUntil.After(now) { if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now) return true, e.lockedUntil.Sub(now)
} }
if !e.lockedUntil.IsZero() { if !e.lockedUntil.IsZero() {
@@ -103,11 +213,14 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
} }
} }
e.fails = append(kept, now) e.fails = append(kept, now)
e.lastFail = now
if len(e.fails) >= pol.maxFails { if len(e.fails) >= pol.maxFails {
e.lockedUntil = now.Add(pol.lockFor) e.lockedUntil = now.Add(pol.lockFor)
e.fails = nil e.fails = nil
l.maybeSweepLocked(now)
return true, pol.lockFor return true, pol.lockFor
} }
l.maybeSweepLocked(now)
return false, 0 return false, 0
} }
@@ -132,6 +245,24 @@ func (l *MemoryLocks) ClearEndpoint(endpointID string) {
} }
} }
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定。
func (l *MemoryLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
if endpointID == "" {
return
}
for k := range l.entries {
parts := splitLockKey(k)
if len(parts) != 4 {
continue
}
if parts[1] == endpointID || parts[2] == endpointID {
delete(l.entries, k)
}
}
}
// Clear 清除精确键。 // Clear 清除精确键。
func (l *MemoryLocks) Clear(key LockKey) { func (l *MemoryLocks) Clear(key LockKey) {
l.mu.Lock() l.mu.Lock()
+60
View File
@@ -0,0 +1,60 @@
package auth
import (
"fmt"
"testing"
"time"
)
func TestLoginLocksSweepExpiredKeepsLocked(t *testing.T) {
locks := NewLoginLocks()
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 10000; i++ {
if locked, _ := locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("2001:db8::%d", i)}); locked {
t.Fatalf("unexpected lock at i=%d", i)
}
}
if n := locks.entryCount(); n < 10000 {
t.Fatalf("want 10000 entries before sweep, got %d", n)
}
now = now.Add(4 * time.Minute)
keep := LockKey{Kind: LockRegisterIP, IP: "keep-locked"}
for i := 0; i < 10; i++ {
locks.Fail(keep)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("keep-locked should be locked")
}
now = now.Add(2 * time.Minute) // 10k 已过 5 分钟窗口;keep 仍锁定至 +5min
if locked, _ := locks.Check(LockKey{Kind: LockRegisterIP, IP: "trigger-sweep"}); locked {
t.Fatal("trigger must not lock")
}
if n := locks.entryCount(); n != 1 {
t.Fatalf("want only locked entry after sweep, got %d", n)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("locked entry must remain")
}
}
func TestLoginLocksCapDoesNotGrow(t *testing.T) {
locks := NewLoginLocks()
locks.maxEntries = 64
now := time.Date(2026, 9, 30, 15, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 200; i++ {
locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("ip-%d", i)})
now = now.Add(time.Millisecond)
if n := locks.entryCount(); n > 64 {
t.Fatalf("entries=%d exceeded cap at i=%d", n, i)
}
}
if n := locks.entryCount(); n != 64 {
t.Fatalf("entries=%d want 64", n)
}
}
+6
View File
@@ -94,6 +94,12 @@ func (l *StubLoginLocks) ClearEndpoint(endpointID string) {
l.Cleared = append(l.Cleared, endpointID) l.Cleared = append(l.Cleared, endpointID)
} }
func (l *StubLoginLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
l.Cleared = append(l.Cleared, endpointID)
}
func (l *StubLoginLocks) Clear(LockKey) {} func (l *StubLoginLocks) Clear(LockKey) {}
// 编译期检查:假实现满足接口。 // 编译期检查:假实现满足接口。
+18 -2
View File
@@ -8,6 +8,17 @@ import (
"net/http" "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 对象。 // ErrorBody 是失败响应里的 error 对象。
type ErrorBody struct { type ErrorBody struct {
Code string `json:"code"` 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 { func DecodeJSON(r *http.Request, dst any) error {
defer func() { _ = r.Body.Close() }() 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() dec.DisallowUnknownFields()
if err := dec.Decode(dst); err != nil { if err := dec.Decode(dst); err != nil {
if errors.Is(err, io.EOF) { if errors.Is(err, io.EOF) {
+39
View File
@@ -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)
}
}
+4 -3
View File
@@ -106,11 +106,12 @@ test.describe("后台主路径 W4", () => {
await nav(page, "注册设置"); await nav(page, "注册设置");
await expect(page).toHaveURL(/\/registration/); await expect(page).toHaveURL(/\/registration/);
await page.getByTestId("reg-enabled").click(); await expect(page.getByTestId("reg-code-missing")).toBeVisible();
await expect(page.getByText("已开启自助注册")).toBeVisible();
await page.getByTestId("reg-code").locator("input").fill("w4-reg-code-01"); await page.getByTestId("reg-code").locator("input").fill("w4-reg-code-01");
await page.getByTestId("reg-save-code").click(); await page.getByTestId("reg-save-code").click();
await expect(page.getByText("安全码已更新")).toBeVisible(); await expect(page.getByText("安全码已更新")).toBeVisible();
await page.getByTestId("reg-enabled").click();
await expect(page.getByText("已开启自助注册")).toBeVisible();
await nav(page, "群"); await nav(page, "群");
await expect(page).toHaveURL(/\/groups/); await expect(page).toHaveURL(/\/groups/);
@@ -126,7 +127,7 @@ test.describe("后台主路径 W4", () => {
await nav(page, "投递记录"); await nav(page, "投递记录");
await expect(page).toHaveURL(/\/messages/); await expect(page).toHaveURL(/\/messages/);
await expect(page.getByText("w4-e2e-scheduled")).toBeVisible(); 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 nav(page, "API 令牌");
await expect(page).toHaveURL(/\/tokens/); await expect(page).toHaveURL(/\/tokens/);
+5 -5
View File
@@ -22,9 +22,7 @@ async function run<T>(fn: () => Promise<T>, silent = false): Promise<T> {
try { try {
return await fn(); return await fn();
} catch (e) { } catch (e) {
if (!silent && e instanceof ApiError) { if (!silent && e instanceof Error && !(e instanceof ApiError)) {
message.error(e.message);
} else if (!silent && e instanceof Error) {
message.error(e.message); message.error(e.message);
} }
throw e; throw e;
@@ -163,12 +161,14 @@ export function listMessages(q: {
endpoint_id?: string; endpoint_id?: string;
group_id?: string; group_id?: string;
state?: string; state?: string;
from_ms?: number;
to_ms?: number;
}) { }) {
return run(() => mockApi.listMessages(q)); return run(() => mockApi.listMessages(q));
} }
export function getMessage(seq: number) { export function getMessage(seq: number, cursor?: string, limit?: number) {
return run(() => mockApi.getMessage(seq)); return run(() => mockApi.getMessage(seq, cursor, limit));
} }
export function getSettings() { export function getSettings() {
+9 -5
View File
@@ -33,9 +33,7 @@ async function run<T>(fn: () => Promise<T>, silent = false): Promise<T> {
try { try {
return await fn(); return await fn();
} catch (e) { } catch (e) {
if (!silent && e instanceof ApiError) { if (!silent && e instanceof Error && !(e instanceof ApiError)) {
message.error(e.message);
} else if (!silent && e instanceof Error) {
message.error(e.message); message.error(e.message);
} }
throw e; throw e;
@@ -333,6 +331,8 @@ export function listMessages(q: {
endpoint_id?: string; endpoint_id?: string;
group_id?: string; group_id?: string;
state?: string; state?: string;
from_ms?: number;
to_ms?: number;
}): Promise<PageResult<MessageSummary>> { }): Promise<PageResult<MessageSummary>> {
return run(() => return run(() =>
requestAdmin<PageResult<MessageSummary>>( requestAdmin<PageResult<MessageSummary>>(
@@ -343,13 +343,17 @@ export function listMessages(q: {
endpoint_id: q.endpoint_id, endpoint_id: q.endpoint_id,
group_id: q.group_id, group_id: q.group_id,
state: q.state, state: q.state,
from_ms: q.from_ms,
to_ms: q.to_ms,
})}`, })}`,
), ),
); );
} }
export function getMessage(seq: number): Promise<MessageDetail> { export function getMessage(seq: number, cursor?: string, limit?: number): Promise<MessageDetail> {
return run(() => requestAdmin<MessageDetail>(`/api/admin/messages/${seq}`)); return run(() =>
requestAdmin<MessageDetail>(`/api/admin/messages/${seq}${buildQuery({ cursor, limit })}`),
);
} }
export function getSettings(): Promise<RuntimeSettings> { export function getSettings(): Promise<RuntimeSettings> {
+78
View File
@@ -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("请求过于频繁,请稍后再试");
});
});
+69 -7
View File
@@ -27,9 +27,52 @@ export interface RequestOptions {
contentType?: string | null; 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 显示。 * 管理接口请求封装:自动带 credentials 与 X-Nixmsg-Request: 1。
* W4 接真实后端时页面不必改,只需让 admin 模块走本函数。 * 错误只在这里提示一次;401(登录/改密除外)清空会话并跳转登录页。
*/ */
export async function requestAdmin<T>(path: string, opts: RequestOptions = {}): Promise<T> { export async function requestAdmin<T>(path: string, opts: RequestOptions = {}): Promise<T> {
const method = opts.method ?? (opts.body != null || opts.rawBody != null ? "POST" : "GET"); const method = opts.method ?? (opts.body != null || opts.rawBody != null ? "POST" : "GET");
@@ -46,19 +89,32 @@ export async function requestAdmin<T>(path: string, opts: RequestOptions = {}):
headers["Content-Type"] = opts.contentType; headers["Content-Type"] = opts.contentType;
} }
const res = await fetch(path, { let res: Response;
try {
res = await fetch(path, {
method, method,
credentials: "include", credentials: "include",
headers, headers,
body: body ?? undefined, 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<T>; let envelope: ApiEnvelope<T>;
try { try {
envelope = (await res.json()) as ApiEnvelope<T>; envelope = (await res.json()) as ApiEnvelope<T>;
} catch { } catch {
const err = new ApiError("internal", `响应不是 JSON(HTTP ${res.status})`, res.status); const err = new ApiError("internal", `响应不是 JSON(HTTP ${res.status})`, res.status);
if (!opts.silent) { if (!opts.silent && res.status !== 401) {
message.error(err.message); message.error(err.message);
} }
throw err; throw err;
@@ -66,7 +122,7 @@ export async function requestAdmin<T>(path: string, opts: RequestOptions = {}):
if (typeof envelope !== "object" || envelope === null) { if (typeof envelope !== "object" || envelope === null) {
const err = new ApiError("internal", "响应格式无效", res.status); const err = new ApiError("internal", "响应格式无效", res.status);
if (!opts.silent) { if (!opts.silent && res.status !== 401) {
message.error(err.message); message.error(err.message);
} }
throw err; throw err;
@@ -77,8 +133,14 @@ export async function requestAdmin<T>(path: string, opts: RequestOptions = {}):
} }
const errBody: ApiErrorBody = envelope.error ?? { code: "internal", message: "未知错误" }; const errBody: ApiErrorBody = envelope.error ?? { code: "internal", message: "未知错误" };
const err = new ApiError(errBody.code, errBody.message, res.status, "data" in envelope ? envelope.data : undefined); let display = errBody.message;
if (!opts.silent) { 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); message.error(err.message);
} }
throw err; throw err;
+25 -5
View File
@@ -151,6 +151,14 @@ const messages: MessageDetail[] = [
pushed_at_ms: now() - 3_599_000, pushed_at_ms: now() - 3_599_000,
updated_at_ms: now() - 3_598_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: "", next_cursor: "",
}, },
@@ -320,6 +328,7 @@ export const mockApi = {
endpoints_total: endpoints.length, endpoints_total: endpoints.length,
endpoints_online: endpoints.filter((e) => e.online).length, endpoints_online: endpoints.filter((e) => e.online).length,
endpoints_disabled: endpoints.filter((e) => !e.enabled).length, endpoints_disabled: endpoints.filter((e) => !e.enabled).length,
endpoints_self: endpoints.filter((e) => e.source === "self").length,
groups_total: groups.length, groups_total: groups.length,
messages_pending: messages.filter((m) => m.state === "dispatched").length, messages_pending: messages.filter((m) => m.state === "dispatched").length,
messages_scheduled: messages.filter((m) => m.state === "scheduled").length, messages_scheduled: messages.filter((m) => m.state === "scheduled").length,
@@ -354,7 +363,8 @@ export const mockApi = {
if (endpoints.some((e) => e.id === id)) { if (endpoints.some((e) => e.id === id)) {
throw new ApiError("id_taken", "编号已占用", 409); 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({ endpoints.unshift({
id, id,
name: body.name, name: body.name,
@@ -369,7 +379,7 @@ export const mockApi = {
created_at_ms: now(), created_at_ms: now(),
login_locked: false, login_locked: false,
}); });
return { id, login_password: loginPassword }; return provided ? { id } : { id, login_password: loginPassword };
}, },
async importEndpoints(csvText: string): Promise<{ items: ImportItem[] }> { async importEndpoints(csvText: string): Promise<{ items: ImportItem[] }> {
@@ -598,7 +608,7 @@ export const mockApi = {
let list = groups.map(toGroupSummary); let list = groups.map(toGroupSummary);
if (query) { if (query) {
const n = query.toLowerCase(); 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); return paginate(list, cursor, limit);
}, },
@@ -655,6 +665,7 @@ export const mockApi = {
name: g.name, name: g.name,
owner_id: g.owner_id, owner_id: g.owner_id,
created_at_ms: g.created_at_ms, created_at_ms: g.created_at_ms,
member_total: all.length,
members: page.items, members: page.items,
next_cursor: page.next_cursor, next_cursor: page.next_cursor,
}; };
@@ -729,6 +740,8 @@ export const mockApi = {
endpoint_id?: string; endpoint_id?: string;
group_id?: string; group_id?: string;
state?: string; state?: string;
from_ms?: number;
to_ms?: number;
}): Promise<PageResult<MessageSummary>> { }): Promise<PageResult<MessageSummary>> {
requireSession(); requireSession();
let list = messages.map(toMessageSummary); 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.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); return paginate(list, q.cursor, q.limit);
}, },
async getMessage(seq: number): Promise<MessageDetail> { async getMessage(seq: number, cursor?: string, limit?: number): Promise<MessageDetail> {
requireSession(); requireSession();
const m = messages.find((x) => x.seq === seq); const m = messages.find((x) => x.seq === seq);
if (!m) { if (!m) {
throw new ApiError("not_found", "消息不存在", 404); 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<RuntimeSettings> { async getSettings(): Promise<RuntimeSettings> {
+4 -4
View File
@@ -34,8 +34,7 @@ export interface Overview {
endpoints_total: number; endpoints_total: number;
endpoints_online: number; endpoints_online: number;
endpoints_disabled: number; endpoints_disabled: number;
/** A3 补充:自助注册端数量;契约示例未列,前端可选展示 */ endpoints_self: number;
endpoints_self?: number;
groups_total: number; groups_total: number;
messages_pending: number; messages_pending: number;
messages_scheduled: number; messages_scheduled: number;
@@ -72,7 +71,7 @@ export interface EndpointCreateRequest {
export interface EndpointCreateResult { export interface EndpointCreateResult {
id: string; id: string;
login_password: string; login_password?: string;
} }
export interface EndpointPatchRequest { export interface EndpointPatchRequest {
@@ -98,7 +97,7 @@ export interface ImportError {
export interface ImportItem { export interface ImportItem {
id: string; id: string;
login_password: string; login_password?: string;
name: string; name: string;
} }
@@ -155,6 +154,7 @@ export interface GroupDetail {
name: string; name: string;
owner_id: string; owner_id: string;
created_at_ms: number; created_at_ms: number;
member_total?: number;
members: GroupMember[]; members: GroupMember[];
next_cursor: string; next_cursor: string;
} }
+19
View File
@@ -0,0 +1,19 @@
<script setup lang="ts">
import { NButton, NResult } from "naive-ui";
defineProps<{
description?: string;
}>();
const emit = defineEmits<{
retry: [];
}>();
</script>
<template>
<n-result status="error" title="加载失败" :description="description || '无法加载页面数据'">
<template #footer>
<n-button type="primary" data-testid="load-retry" @click="emit('retry')">重试</n-button>
</template>
</n-result>
</template>
@@ -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();
});
});
+46 -15
View File
@@ -1,40 +1,71 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref } from "vue";
import { NAlert, NButton, NSpace, NText } from "naive-ui"; import { NAlert, NButton, NSpace, NText } from "naive-ui";
import { downloadText } from "@/utils/csv";
import { message } from "@/utils/notify";
const props = defineProps<{ const props = withDefaults(
defineProps<{
title: string; title: string;
secret: string; secret: string;
filename?: string; filename?: string;
}>(); mime?: string;
}>(),
{
filename: "secret.txt",
mime: "text/plain;charset=utf-8",
},
);
const emit = defineEmits<{ const emit = defineEmits<{
dismiss: []; dismiss: [];
}>(); }>();
function copy() { const secretEl = ref<HTMLElement | null>(null);
void navigator.clipboard?.writeText(props.secret);
async function copy() {
try {
if (typeof navigator !== "undefined" && navigator.clipboard?.writeText) {
await navigator.clipboard.writeText(props.secret);
message.success("已复制");
return;
}
} catch {
/* 回退选中文本 */
}
const el = secretEl.value;
if (el) {
const sel = window.getSelection();
const range = document.createRange();
range.selectNodeContents(el);
sel?.removeAllRanges();
sel?.addRange(range);
}
try {
if (document.execCommand("copy")) {
message.success("已复制");
return;
}
} catch {
/* ignore */
}
message.error("复制失败,请手动选择文本");
} }
function download() { function download() {
const blob = new Blob([props.secret], { type: "text/plain;charset=utf-8" }); downloadText(props.filename, props.secret, props.mime);
const url = URL.createObjectURL(blob);
const a = document.createElement("a");
a.href = url;
a.download = props.filename ?? "secret.txt";
a.click();
URL.revokeObjectURL(url);
} }
</script> </script>
<template> <template>
<n-alert type="warning" :title="title" style="margin-bottom: 12px"> <n-alert type="warning" :title="title" style="margin-bottom: 12px">
<n-space vertical :size="8"> <n-space vertical :size="8">
<n-text code>{{ secret }}</n-text> <span ref="secretEl" data-testid="secret-text"><n-text code>{{ secret }}</n-text></span>
<n-text depth="3">关闭后将无法再次查看,请立即复制或下载。</n-text> <n-text depth="3">关闭后将无法再次查看,请立即复制或下载。</n-text>
<n-space> <n-space>
<n-button size="small" @click="copy">复制</n-button> <n-button size="small" data-testid="secret-copy" @click="copy">复制</n-button>
<n-button size="small" @click="download">下载</n-button> <n-button size="small" data-testid="secret-download" @click="download">下载</n-button>
<n-button size="small" secondary @click="emit('dismiss')">我已保存</n-button> <n-button size="small" secondary data-testid="secret-dismiss" @click="emit('dismiss')">我已保存</n-button>
</n-space> </n-space>
</n-space> </n-space>
</n-alert> </n-alert>
+15 -1
View File
@@ -2,8 +2,22 @@ import { createApp } from "vue";
import { createPinia } from "pinia"; import { createPinia } from "pinia";
import App from "./App.vue"; import App from "./App.vue";
import { router } from "./router"; import { router } from "./router";
import { setUnauthorizedHandler } from "./api/http";
import { useAuthStore } from "./stores/auth";
const app = createApp(App); const app = createApp(App);
app.use(createPinia()); const pinia = createPinia();
app.use(pinia);
app.use(router); app.use(router);
setUnauthorizedHandler((redirect) => {
const auth = useAuthStore();
auth.clearSession();
if (router.currentRoute.value.meta.public) {
return;
}
const target = redirect.startsWith("/login") ? "/overview" : redirect || "/overview";
void router.replace({ name: "login", query: { redirect: target } });
});
app.mount("#app"); app.mount("#app");
+5 -1
View File
@@ -32,5 +32,9 @@ export const useAuthStore = defineStore("auth", () => {
} }
} }
return { username, ready, isLoggedIn, hydrate, login, logout }; function clearSession() {
username.value = null;
}
return { username, ready, isLoggedIn, hydrate, login, logout, clearSession };
}); });
+36
View File
@@ -0,0 +1,36 @@
import { describe, expect, it } from "vitest";
import { buildPasswordCsv, csvQuote, guardFormula } from "./csv";
describe("csvQuote", () => {
it("给所有字段加引号并双写内部引号", () => {
expect(csvQuote("a")).toBe('"a"');
expect(csvQuote('say "hi"')).toBe('"say ""hi"""');
expect(csvQuote("a,b")).toBe('"a,b"');
expect(csvQuote("a\nb")).toBe('"a\nb"');
});
});
describe("guardFormula", () => {
it("名称以公式字符开头时前缀单引号", () => {
expect(guardFormula("=1+1")).toBe("'=1+1");
expect(guardFormula("+cmd")).toBe("'+cmd");
expect(guardFormula("-1")).toBe("'-1");
expect(guardFormula("@SUM")).toBe("'@SUM");
expect(guardFormula("门口")).toBe("门口");
});
});
describe("buildPasswordCsv", () => {
it("带 BOM、表头、CRLF,密码原样加引号", () => {
const out = buildPasswordCsv([
{ id: "e1", name: "门,口", login_password: 'p"w' },
{ id: "e2", name: "=cmd", login_password: "=1+1" },
]);
expect(out.startsWith("\uFEFF")).toBe(true);
const body = out.slice(1);
expect(body.startsWith('"id","name","login_password"\r\n')).toBe(true);
const rows = body.split("\r\n");
expect(rows[1]).toBe('"e1","门,口","p""w"');
expect(rows[2]).toBe('"e2","\'=cmd","=1+1"');
});
});
+40
View File
@@ -0,0 +1,40 @@
/** RFC 4180 CSV,用于一次性密码下载。 */
const FORMULA_START = /^[=+\-@]/;
export function csvQuote(value: string): string {
return `"${String(value ?? "").replace(/"/g, '""')}"`;
}
/** 名称列以防公式注入:以 = + - @ 开头时前缀单引号。密码列不要调用。 */
export function guardFormula(value: string): string {
const v = String(value ?? "");
return FORMULA_START.test(v) ? `'${v}` : v;
}
export interface PasswordCsvRow {
id: string;
name: string;
login_password: string;
}
/** BOM + 表头 + CRLF;所有字段加引号。 */
export function buildPasswordCsv(rows: PasswordCsvRow[]): string {
const lines = [
[csvQuote("id"), csvQuote("name"), csvQuote("login_password")].join(","),
...rows.map((r) =>
[csvQuote(r.id), csvQuote(guardFormula(r.name)), csvQuote(r.login_password)].join(","),
),
];
return `\uFEFF${lines.join("\r\n")}`;
}
export function downloadText(filename: string, text: string, mime: string) {
const blob = new Blob([text], { type: mime });
const url = URL.createObjectURL(blob);
const a = document.createElement("a");
a.href = url;
a.download = filename;
a.click();
URL.revokeObjectURL(url);
}
+15
View File
@@ -0,0 +1,15 @@
import { describe, expect, it } from "vitest";
import { formatUptime, zhLabel, messageStateLabel } from "./labels";
describe("labels", () => {
it("消息状态译成中文", () => {
expect(zhLabel(messageStateLabel, "scheduled")).toBe("定时中");
expect(zhLabel(messageStateLabel, "unknown")).toBe("unknown");
});
it("运行时长按天时分", () => {
expect(formatUptime(90_000)).toBe("1 分");
expect(formatUptime(3_661_000)).toBe("1 小时 1 分");
expect(formatUptime(90_000_000)).toBe("1 天 1 小时 0 分");
});
});
+75
View File
@@ -0,0 +1,75 @@
/** 后台界面中文标签;原值可放 tooltip。 */
export const messageStateLabel: Record<string, string> = {
scheduled: "定时中",
dispatched: "投递中",
completed: "已完成",
};
export const destKindLabel: Record<string, string> = {
endpoint: "端",
group: "群",
};
export const deliveryStateLabel: Record<string, string> = {
pending: "待投递",
accepted: "已收下",
recalled: "已撤回",
expired: "已过期",
dropped: "已丢弃",
rejected: "已拒绝",
};
export const reasonLabel: Record<string, string> = {
disabled: "端已停用",
expired: "已过期",
recalled: "已撤回",
dropped: "已丢弃",
rejected: "已拒绝",
unauthorized: "未授权",
not_found: "不存在",
};
export const settingsFieldLabel: Record<string, string> = {
listen: "端监听地址",
admin_listen: "后台监听地址",
session_idle_days: "会话闲置天数",
record_retention_days: "记录保留天数",
idempotency_hours: "防重小时数",
receipt_retention_days: "回执保留天数",
sqlite_synchronous: "SQLite 同步模式",
};
export const limitsLabel: Record<string, string> = {
max_body_bytes: "正文上限(字节)",
max_meta_bytes: "元数据上限(字节)",
max_frame_bytes: "帧上限(字节)",
max_ttl_seconds: "TTL 上限(秒)",
max_schedule_seconds: "定时上限(秒)",
max_group_members: "群成员上限",
grace_seconds: "宽限秒数",
ack_timeout_seconds: "确认超时(秒)",
delivery_window: "投递窗口",
receipt_window: "回执窗口",
requests_per_second: "每秒请求",
max_pending_per_sender: "发送方未完成上限",
max_pending_per_receiver: "接收方未完成上限",
};
export function zhLabel(map: Record<string, string>, value: string | undefined | null, fallback = "—"): string {
if (value == null || value === "") return fallback;
return map[value] ?? value;
}
/** 运行时长:天、时、分。 */
export function formatUptime(ms: number): string {
const total = Math.max(0, Math.floor(ms / 1000));
const days = Math.floor(total / 86400);
const hours = Math.floor((total % 86400) / 3600);
const mins = Math.floor((total % 3600) / 60);
const parts: string[] = [];
if (days > 0) parts.push(`${days} 天`);
if (hours > 0 || days > 0) parts.push(`${hours} 小时`);
parts.push(`${mins} 分`);
return parts.join(" ");
}
+104
View File
@@ -11,6 +11,7 @@ import {
import { defineComponent, h, nextTick } from "vue"; import { defineComponent, h, nextTick } from "vue";
import EndpointsView from "./EndpointsView.vue"; import EndpointsView from "./EndpointsView.vue";
import { mockApi } from "@/api/mock"; import { mockApi } from "@/api/mock";
import * as admin from "@/api/admin";
vi.mock("@/api/admin", async () => import("@/api/admin-mock")); vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
@@ -69,4 +70,107 @@ describe("EndpointsView 开通端", () => {
expect(w.text()).not.toContain("只显示一次"); expect(w.text()).not.toContain("只显示一次");
w.unmount(); w.unmount();
}); });
it("手填密码开通后不显示一次性面板和 undefined", async () => {
const pinia = createPinia();
const w = mount(wrap(EndpointsView), {
global: { plugins: [pinia] },
attachTo: document.body,
});
await flushPromises();
const openBtn = w.findAll("button").find((b) => b.text().trim() === "开通");
await openBtn!.trigger("click");
await nextTick();
await flushPromises();
await w.find('[data-testid="create-name"]').find("input").setValue("手填端");
await w.find('[data-testid="create-login-password"]').find("input").setValue("handpassword12");
await w.find('[data-testid="create-submit"]').trigger("click");
await flushPromises();
expect(w.text()).not.toContain("undefined");
expect(w.text()).not.toContain("只显示一次");
expect(w.text()).toContain("手填端");
w.unmount();
});
});
describe("EndpointsView 列表与编辑", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
vi.restoreAllMocks();
});
it("login_locked=false 时解锁按钮可点", async () => {
const w = mount(wrap(EndpointsView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
const unlocks = w.findAll('[data-testid="endpoint-unlock"]');
expect(unlocks.length).toBeGreaterThan(0);
for (const b of unlocks) {
expect(b.attributes("disabled")).toBeUndefined();
}
w.unmount();
});
it("只改名称时请求体没有 enabled", async () => {
const spy = vi.spyOn(admin, "patchEndpoint");
const w = mount(wrap(EndpointsView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
const editBtn = w.findAll("button").find((b) => b.text().trim() === "编辑");
await editBtn!.trigger("click");
await nextTick();
await flushPromises();
await w.find('[data-testid="edit-name"]').find("input").setValue("改名后");
await w.find('[data-testid="edit-submit"]').trigger("click");
await flushPromises();
expect(spy).toHaveBeenCalled();
const body = spy.mock.calls[0][1];
expect(body).toEqual({ name: "改名后" });
expect(body).not.toHaveProperty("enabled");
w.unmount();
});
it("默认延迟超上限时前端阻止提交", async () => {
vi.spyOn(admin, "getSettings").mockResolvedValue({
listen: ":7443",
admin_listen: "",
session_idle_days: 30,
record_retention_days: 7,
idempotency_hours: 24,
receipt_retention_days: 7,
sqlite_synchronous: "FULL",
limits: { max_schedule_seconds: 10 },
});
const createSpy = vi.spyOn(admin, "createEndpoint");
const w = mount(wrap(EndpointsView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
const openBtn = w.findAll("button").find((b) => b.text().trim() === "开通");
await openBtn!.trigger("click");
await nextTick();
await flushPromises();
expect(w.find('[data-testid="create-delay-max"]').text()).toContain("上限 10 秒");
const delay = w.find('[data-testid="create-delay"]').find("input");
await delay.setValue("11");
await w.find('[data-testid="create-submit"]').trigger("click");
await flushPromises();
expect(createSpy.mock.calls.every((c) => ((c[0] as { default_delay_seconds?: number }).default_delay_seconds ?? 0) <= 10)).toBe(true);
w.unmount();
});
}); });
+249 -40
View File
@@ -2,6 +2,7 @@
import { computed, h, onMounted, reactive, ref } from "vue"; import { computed, h, onMounted, reactive, ref } from "vue";
import type { DataTableColumns, DataTableRowKey } from "naive-ui"; import type { DataTableColumns, DataTableRowKey } from "naive-ui";
import { import {
NAlert,
NButton, NButton,
NDataTable, NDataTable,
NForm, NForm,
@@ -21,9 +22,11 @@ import type { UploadCustomRequestOptions } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import HelpTip from "@/components/HelpTip.vue"; import HelpTip from "@/components/HelpTip.vue";
import SecretOnceAlert from "@/components/SecretOnceAlert.vue"; import SecretOnceAlert from "@/components/SecretOnceAlert.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import { import {
batchEndpoints, batchEndpoints,
createEndpoint, createEndpoint,
getSettings,
importEndpoints, importEndpoints,
kickEndpoint, kickEndpoint,
listEndpoints, listEndpoints,
@@ -32,15 +35,18 @@ import {
setTalkPassword, setTalkPassword,
unlockEndpoint, unlockEndpoint,
type Endpoint, type Endpoint,
type EndpointPatchRequest,
type ImportItem, type ImportItem,
} from "@/api/admin"; } from "@/api/admin";
import { ApiError } from "@/api/http"; import { ApiError } from "@/api/http";
import { formatLocalMs } from "@/utils/time"; import { formatLocalMs } from "@/utils/time";
import { buildPasswordCsv, downloadText } from "@/utils/csv";
import { message } from "@/utils/notify"; import { message } from "@/utils/notify";
const dialog = useDialog(); const dialog = useDialog();
const loading = ref(false); const loading = ref(false);
const loadError = ref("");
const rows = ref<Endpoint[]>([]); const rows = ref<Endpoint[]>([]);
const total = ref(0); const total = ref(0);
const checkedKeys = ref<DataTableRowKey[]>([]); const checkedKeys = ref<DataTableRowKey[]>([]);
@@ -64,8 +70,12 @@ const onlineOptions = [
{ label: "离线", value: "false" }, { label: "离线", value: "false" },
]; ];
const onceSecret = ref<{ title: string; secret: string; filename: string } | null>(null); const onceSecret = ref<{ title: string; secret: string; filename: string; mime?: string } | null>(null);
const importResult = ref<ImportItem[] | null>(null); const importResult = ref<ImportItem[] | null>(null);
const importing = ref(false);
const importEta = ref(1);
const importErrors = ref<{ line: number; reason: string }[] | null>(null);
const maxDelaySeconds = ref(31536000);
const createOpen = ref(false); const createOpen = ref(false);
const createForm = reactive({ const createForm = reactive({
@@ -77,6 +87,7 @@ const createForm = reactive({
default_delay_seconds: 0, default_delay_seconds: 0,
}); });
const createLoading = ref(false); const createLoading = ref(false);
const createError = ref("");
const editOpen = ref(false); const editOpen = ref(false);
const editId = ref(""); const editId = ref("");
@@ -87,6 +98,12 @@ const editForm = reactive({
enabled: true, enabled: true,
}); });
const editLoading = ref(false); const editLoading = ref(false);
const editOrig = reactive({
name: "",
remark: "",
default_delay_seconds: 0,
enabled: true,
});
const talkOpen = ref(false); const talkOpen = ref(false);
const talkId = ref(""); const talkId = ref("");
@@ -94,6 +111,7 @@ const talkPassword = ref("");
async function load() { async function load() {
loading.value = true; loading.value = true;
loadError.value = "";
try { try {
const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : ""; const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : "";
const res = await listEndpoints({ const res = await listEndpoints({
@@ -105,6 +123,8 @@ async function load() {
}); });
rows.value = res.items; rows.value = res.items;
total.value = res.total; total.value = res.total;
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
@@ -112,8 +132,21 @@ async function load() {
onMounted(() => { onMounted(() => {
void load(); void load();
void loadLimits();
}); });
async function loadLimits() {
try {
const s = await getSettings();
const n = s.limits?.max_schedule_seconds;
if (typeof n === "number" && n > 0) {
maxDelaySeconds.value = n;
}
} catch {
/* 沿用默认 365 天 */
}
}
function sourceLabel(s: string) { function sourceLabel(s: string) {
return s === "self" ? "自助注册" : "后台开通"; return s === "self" ? "自助注册" : "后台开通";
} }
@@ -141,16 +174,29 @@ const columns = computed<DataTableColumns<Endpoint>>(() => [
render: (r) => (r.enabled ? "是" : "否"), render: (r) => (r.enabled ? "是" : "否"),
}, },
{ {
title: "锁定", title() {
return h("span", [
"锁定",
h(HelpTip, null, {
default: () => "列表只显示按编号的总数锁定,解锁会同时清除按 IP 的锁定",
}),
]);
},
key: "login_locked", key: "login_locked",
width: 70, width: 90,
render: (r) => (r.login_locked ? "是" : "否"), render: (r) => (r.login_locked ? "是" : "否"),
}, },
{ {
title: "最近在线", title: "最近上线",
key: "online_since_ms", key: "online_since_ms",
width: 160, width: 160,
render: (r) => formatLocalMs(r.online ? r.online_since_ms : r.offline_since_ms), render: (r) => formatLocalMs(r.online_since_ms),
},
{
title: "最近离线",
key: "offline_since_ms",
width: 160,
render: (r) => formatLocalMs(r.offline_since_ms),
}, },
{ {
title: "创建时间", title: "创建时间",
@@ -170,7 +216,11 @@ const columns = computed<DataTableColumns<Endpoint>>(() => [
h(NButton, { size: "tiny", onClick: () => openTalk(r) }, () => "对话密码"), h(NButton, { size: "tiny", onClick: () => openTalk(r) }, () => "对话密码"),
h( h(
NButton, NButton,
{ size: "tiny", disabled: !r.login_locked, onClick: () => void onUnlock(r) }, {
size: "tiny",
"data-testid": "endpoint-unlock",
onClick: () => void onUnlock(r),
},
() => "解锁", () => "解锁",
), ),
h(NButton, { size: "tiny", onClick: () => void onKick(r) }, () => "踢下线"), h(NButton, { size: "tiny", onClick: () => void onKick(r) }, () => "踢下线"),
@@ -187,12 +237,14 @@ function openCreate() {
talk_password: "", talk_password: "",
default_delay_seconds: 0, default_delay_seconds: 0,
}); });
createError.value = "";
createOpen.value = true; createOpen.value = true;
} }
async function submitCreate() { async function submitCreate() {
if (!createForm.name.trim()) { createError.value = "";
message.error("名称不能为空"); if ((createForm.default_delay_seconds ?? 0) > maxDelaySeconds.value) {
createError.value = `默认延迟不能超过 ${maxDelaySeconds.value} 秒`;
return; return;
} }
createLoading.value = true; createLoading.value = true;
@@ -206,12 +258,18 @@ async function submitCreate() {
default_delay_seconds: createForm.default_delay_seconds, default_delay_seconds: createForm.default_delay_seconds,
}); });
createOpen.value = false; createOpen.value = false;
if (res.login_password) {
onceSecret.value = { onceSecret.value = {
title: "开通成功,登录密码只显示一次", title: "开通成功,登录密码只显示一次",
secret: `${res.id}\t${res.login_password}`, secret: `${res.id}\t${res.login_password}`,
filename: `${res.id}-password.txt`, filename: `${res.id}-password.txt`,
}; };
} else {
message.success(`已开通 ${res.id}`);
}
await load(); await load();
} catch (e) {
createError.value = e instanceof ApiError ? e.message : e instanceof Error ? e.message : "开通失败";
} finally { } finally {
createLoading.value = false; createLoading.value = false;
} }
@@ -225,24 +283,51 @@ function openEdit(r: Endpoint) {
default_delay_seconds: Math.round(r.default_delay_ms / 1000), default_delay_seconds: Math.round(r.default_delay_ms / 1000),
enabled: r.enabled, enabled: r.enabled,
}); });
Object.assign(editOrig, editForm);
editOpen.value = true; editOpen.value = true;
} }
function buildEditPatch(): EndpointPatchRequest {
const body: EndpointPatchRequest = {};
if (editForm.name !== editOrig.name) body.name = editForm.name;
if (editForm.remark !== editOrig.remark) body.remark = editForm.remark;
if (editForm.default_delay_seconds !== editOrig.default_delay_seconds) {
body.default_delay_seconds = editForm.default_delay_seconds;
}
if (editForm.enabled !== editOrig.enabled) body.enabled = editForm.enabled;
return body;
}
async function submitEdit() { async function submitEdit() {
if ((editForm.default_delay_seconds ?? 0) > maxDelaySeconds.value) {
message.error(`默认延迟不能超过 ${maxDelaySeconds.value} 秒`);
return;
}
const body = buildEditPatch();
const doSave = async () => {
editLoading.value = true; editLoading.value = true;
try { try {
await patchEndpoint(editId.value, { if (Object.keys(body).length) {
name: editForm.name, await patchEndpoint(editId.value, body);
remark: editForm.remark, }
default_delay_seconds: editForm.default_delay_seconds,
enabled: editForm.enabled,
});
editOpen.value = false; editOpen.value = false;
message.success("已保存"); message.success("已保存");
await load(); await load();
} finally { } finally {
editLoading.value = false; editLoading.value = false;
} }
};
if (editOrig.enabled && !editForm.enabled) {
dialog.warning({
title: "停用端",
content: "停用会立即断开并作废未送达消息,不可恢复",
positiveText: "停用",
negativeText: "取消",
onPositiveClick: () => doSave(),
});
return;
}
await doSave();
} }
function openTalk(r: Endpoint) { function openTalk(r: Endpoint) {
@@ -252,10 +337,23 @@ function openTalk(r: Endpoint) {
} }
async function submitTalk() { async function submitTalk() {
const save = async () => {
await setTalkPassword(talkId.value, talkPassword.value); await setTalkPassword(talkId.value, talkPassword.value);
talkOpen.value = false; talkOpen.value = false;
message.success(talkPassword.value ? "已设置对话密码" : "已清除对话密码"); message.success(talkPassword.value ? "已设置对话密码" : "已清除对话密码");
await load(); await load();
};
if (!talkPassword.value) {
dialog.warning({
title: "清除对话密码",
content: "留空保存将清除该端的对话密码,确定继续?",
positiveText: "清除",
negativeText: "取消",
onPositiveClick: () => save(),
});
return;
}
await save();
} }
async function onResetPwd(r: Endpoint) { async function onResetPwd(r: Endpoint) {
@@ -293,14 +391,28 @@ async function onBatch(action: "disable" | "delete") {
message.warning("请先多选"); message.warning("请先多选");
return; return;
} }
const preview = ids.slice(0, 20).join("、");
const countPart = ids.length > 20 ? `…共 ${ids.length} 个` : `(${ids.length} 个)`;
let content = `编号:${preview}${countPart}`;
if (action === "delete") {
content += "。删除将退群、必要时转让或解散群主,并清理相关记录,不可恢复。";
} else {
content += "。停用会立即断开并作废未送达消息,不可恢复。";
}
dialog.warning({ dialog.warning({
title: action === "disable" ? "批量停用" : "批量删除", title: action === "disable" ? "批量停用" : "批量删除",
content: `对 ${ids.length} 个端执行${action === "disable" ? "停用" : "删除"}?`, content,
positiveText: "确定", positiveText: "确定",
negativeText: "取消", negativeText: "取消",
onPositiveClick: async () => { onPositiveClick: async () => {
const res = await batchEndpoints(ids, action); const res = await batchEndpoints(ids, action);
message.success(`成功 ${res.ok_ids.length},失败 ${res.failed.length}`); if (res.failed.length) {
message.warning(
`成功 ${res.ok_ids.length}。失败:${res.failed.map((f) => `${f.id}(${f.code})`).join("、")}`,
);
} else {
message.success(`成功 ${res.ok_ids.length}`);
}
checkedKeys.value = []; checkedKeys.value = [];
await load(); await load();
}, },
@@ -313,34 +425,62 @@ async function onImport(options: UploadCustomRequestOptions) {
options.onError(); options.onError();
return; return;
} }
importing.value = true;
importErrors.value = null;
try { try {
const text = await raw.text(); const text = await raw.text();
const lineCount = Math.max(0, text.replace(/^\uFEFF/, "").split(/\r?\n/).filter((l) => l.trim()).length - 1);
const workers = Math.max(1, (navigator.hardwareConcurrency || 4) - 1);
importEta.value = Math.max(1, Math.ceil((lineCount * 0.3) / workers));
const res = await importEndpoints(text); const res = await importEndpoints(text);
importResult.value = res.items; importResult.value = res.items;
const pwdRows = res.items.filter((i): i is ImportItem & { login_password: string } => Boolean(i.login_password));
if (pwdRows.length) {
onceSecret.value = { onceSecret.value = {
title: "CSV 导入成功,密码只显示一次", title: "CSV 导入成功,密码只显示一次",
secret: res.items.map((i) => `${i.id},${i.name},${i.login_password}`).join("\n"), secret: buildPasswordCsv(pwdRows),
filename: "import-passwords.csv", filename: "import-passwords.csv",
mime: "text/csv;charset=utf-8",
}; };
}
options.onFinish(); options.onFinish();
await load(); await load();
} catch (e) { } catch (e) {
options.onError(); options.onError();
if (e instanceof ApiError && e.data && typeof e.data === "object" && "errors" in (e.data as object)) { if (e instanceof ApiError && e.data && typeof e.data === "object" && "errors" in (e.data as object)) {
const errors = (e.data as { errors: { line: number; reason: string }[] }).errors; importErrors.value = (e.data as { errors: { line: number; reason: string }[] }).errors;
message.error(errors.map((x) => `第 ${x.line} 行:${x.reason}`).join(";") || e.message);
} else if (e instanceof Error) { } else if (e instanceof Error) {
message.error(e.message); message.error(e.message);
} }
} finally {
importing.value = false;
} }
} }
function downloadImportCsv() {
if (!importResult.value?.length) {
return;
}
const pwdRows = importResult.value.filter((i): i is ImportItem & { login_password: string } =>
Boolean(i.login_password),
);
downloadText("import-passwords.csv", buildPasswordCsv(pwdRows), "text/csv;charset=utf-8");
}
function reloadFirst() {
checkedKeys.value = [];
page.value = 1;
void load();
}
function onPageChange(p: number) { function onPageChange(p: number) {
checkedKeys.value = [];
page.value = p; page.value = p;
void load(); void load();
} }
function onPageSizeChange(s: number) { function onPageSizeChange(s: number) {
checkedKeys.value = [];
pageSize.value = s; pageSize.value = s;
page.value = 1; page.value = 1;
void load(); void load();
@@ -352,46 +492,59 @@ function onPageSizeChange(s: number) {
<PageHeader title="端" help="可按来源与在线筛选;开通/重置/导入生成的密码只显示一次。"> <PageHeader title="端" help="可按来源与在线筛选;开通/重置/导入生成的密码只显示一次。">
<template #actions> <template #actions>
<n-button type="primary" @click="openCreate">开通</n-button> <n-button type="primary" @click="openCreate">开通</n-button>
<n-upload :show-file-list="false" accept=".csv,text/csv" :custom-request="onImport"> <n-upload :show-file-list="false" accept=".csv,text/csv" :disabled="importing" :custom-request="onImport">
<n-button>CSV 导入</n-button> <n-button :disabled="importing" data-testid="endpoint-import">CSV 导入</n-button>
</n-upload> </n-upload>
<n-button @click="onBatch('disable')">批量停用</n-button> <n-button @click="onBatch('disable')">批量停用</n-button>
<n-button type="error" secondary @click="onBatch('delete')">批量删除</n-button> <n-button type="error" secondary @click="onBatch('delete')">批量删除</n-button>
</template> </template>
</PageHeader> </PageHeader>
<div class="page-body"> <div class="page-body">
<LoadFailed v-if="loadError" :description="loadError" @retry="load" />
<template v-else>
<div class="page-filters">
<SecretOnceAlert <SecretOnceAlert
v-if="onceSecret" v-if="onceSecret"
:title="onceSecret.title" :title="onceSecret.title"
:secret="onceSecret.secret" :secret="onceSecret.secret"
:filename="onceSecret.filename" :filename="onceSecret.filename"
:mime="onceSecret.mime"
@dismiss="onceSecret = null" @dismiss="onceSecret = null"
/> />
<n-alert
v-if="importing"
type="info"
:title="`正在校验并计算密码,约需 ${importEta} 秒`"
style="margin-bottom: 12px"
data-testid="import-progress"
/>
<n-space style="margin-bottom: 12px" align="center"> <n-space style="margin-bottom: 12px" align="center">
<n-select v-model:value="filters.source" :options="sourceOptions" style="width: 140px" @update:value="page = 1; load()" /> <n-select v-model:value="filters.source" :options="sourceOptions" style="width: 140px" @update:value="reloadFirst" />
<n-select v-model:value="filters.online" :options="onlineOptions" style="width: 120px" @update:value="page = 1; load()" /> <n-select v-model:value="filters.online" :options="onlineOptions" style="width: 120px" @update:value="reloadFirst" />
<n-input <n-input
v-model:value="filters.query" v-model:value="filters.query"
clearable clearable
placeholder="编号/名称" placeholder="编号/名称"
style="width: 180px" style="width: 180px"
@keyup.enter="page = 1; load()" @keyup.enter="reloadFirst"
/> />
<n-button @click="page = 1; load()">查询</n-button> <n-button @click="reloadFirst">查询</n-button>
<span> <span>
来源筛选 来源筛选
<HelpTip>admin=后台开通,self=自助注册。</HelpTip> <HelpTip>admin=后台开通,self=自助注册。</HelpTip>
</span> </span>
</n-space> </n-space>
</div>
<div class="page-table">
<n-data-table <n-data-table
remote remote
flex-height
:loading="loading" :loading="loading"
:columns="columns" :columns="columns"
:data="rows" :data="rows"
:row-key="(r: Endpoint) => r.id" :row-key="(r: Endpoint) => r.id"
:checked-row-keys="checkedKeys" :checked-row-keys="checkedKeys"
:max-height="480" :scroll-x="1400"
:scroll-x="1200"
:pagination="{ :pagination="{
page, page,
pageSize, pageSize,
@@ -404,24 +557,38 @@ function onPageSizeChange(s: number) {
@update:checked-row-keys="(k) => (checkedKeys = k)" @update:checked-row-keys="(k) => (checkedKeys = k)"
/> />
</div> </div>
</template>
</div>
<n-modal v-model:show="createOpen" preset="card" title="开通端" style="width: 520px" :mask-closable="false"> <n-modal v-model:show="createOpen" preset="card" title="开通端" style="width: 520px" :mask-closable="false">
<n-scrollbar style="max-height: 360px"> <n-scrollbar style="max-height: 360px">
<n-form label-placement="left" label-width="110" :show-feedback="false"> <n-alert v-if="createError" type="error" :title="createError" style="margin-bottom: 8px" />
<n-form label-placement="left" label-width="110">
<n-form-item label="编号"> <n-form-item label="编号">
<n-input v-model:value="createForm.id" placeholder="留空则生成" data-testid="create-id" /> <n-input v-model:value="createForm.id" placeholder="留空则生成" data-testid="create-id" />
</n-form-item> </n-form-item>
<n-form-item label="名称"> <n-form-item label="名称">
<n-input v-model:value="createForm.name" data-testid="create-name" /> <n-input v-model:value="createForm.name" data-testid="create-name" placeholder="可选" />
</n-form-item> </n-form-item>
<n-form-item label="登录密码"> <n-form-item label="登录密码">
<n-input v-model:value="createForm.login_password" placeholder="留空则生成" type="password" show-password-on="click" /> <n-input v-model:value="createForm.login_password" placeholder="留空则生成" type="password" show-password-on="click" data-testid="create-login-password" />
</n-form-item> </n-form-item>
<n-form-item label="对话密码"> <n-form-item label="对话密码">
<n-input v-model:value="createForm.talk_password" placeholder="可选" type="password" show-password-on="click" /> <n-input v-model:value="createForm.talk_password" placeholder="可选" type="password" show-password-on="click" />
</n-form-item> </n-form-item>
<n-form-item label="默认延迟(秒)"> <n-form-item>
<n-input-number v-model:value="createForm.default_delay_seconds" :min="0" style="width: 100%" /> <template #label>
默认延迟(秒)
<HelpTip>不超过调度上限 {{ maxDelaySeconds }} 秒</HelpTip>
</template>
<n-input-number
v-model:value="createForm.default_delay_seconds"
:min="0"
:max="maxDelaySeconds"
style="width: 100%"
data-testid="create-delay"
/>
<n-text depth="3" data-testid="create-delay-max">上限 {{ maxDelaySeconds }} 秒</n-text>
</n-form-item> </n-form-item>
<n-form-item label="备注" label-placement="top"> <n-form-item label="备注" label-placement="top">
<n-input v-model:value="createForm.remark" type="textarea" :rows="3" /> <n-input v-model:value="createForm.remark" type="textarea" :rows="3" />
@@ -440,13 +607,22 @@ function onPageSizeChange(s: number) {
<n-scrollbar style="max-height: 320px"> <n-scrollbar style="max-height: 320px">
<n-form label-placement="left" label-width="110" :show-feedback="false"> <n-form label-placement="left" label-width="110" :show-feedback="false">
<n-form-item label="名称"> <n-form-item label="名称">
<n-input v-model:value="editForm.name" /> <n-input v-model:value="editForm.name" data-testid="edit-name" />
</n-form-item> </n-form-item>
<n-form-item label="默认延迟(秒)"> <n-form-item>
<n-input-number v-model:value="editForm.default_delay_seconds" :min="0" style="width: 100%" /> <template #label>
默认延迟(秒)
<HelpTip>不超过调度上限 {{ maxDelaySeconds }} 秒</HelpTip>
</template>
<n-input-number
v-model:value="editForm.default_delay_seconds"
:min="0"
:max="maxDelaySeconds"
style="width: 100%"
/>
</n-form-item> </n-form-item>
<n-form-item label="启用"> <n-form-item label="启用">
<n-switch v-model:value="editForm.enabled" /> <n-switch v-model:value="editForm.enabled" data-testid="edit-enabled" />
</n-form-item> </n-form-item>
<n-form-item label="备注" label-placement="top"> <n-form-item label="备注" label-placement="top">
<n-input v-model:value="editForm.remark" type="textarea" :rows="3" /> <n-input v-model:value="editForm.remark" type="textarea" :rows="3" />
@@ -456,7 +632,7 @@ function onPageSizeChange(s: number) {
<template #footer> <template #footer>
<n-space justify="end"> <n-space justify="end">
<n-button @click="editOpen = false">取消</n-button> <n-button @click="editOpen = false">取消</n-button>
<n-button type="primary" :loading="editLoading" @click="submitEdit">保存</n-button> <n-button type="primary" :loading="editLoading" data-testid="edit-submit" @click="submitEdit">保存</n-button>
</n-space> </n-space>
</template> </template>
</n-modal> </n-modal>
@@ -464,7 +640,7 @@ function onPageSizeChange(s: number) {
<n-modal v-model:show="talkOpen" preset="card" title="设置对话密码" style="width: 420px"> <n-modal v-model:show="talkOpen" preset="card" title="设置对话密码" style="width: 420px">
<n-form label-placement="left" label-width="90" :show-feedback="false"> <n-form label-placement="left" label-width="90" :show-feedback="false">
<n-form-item label="对话密码"> <n-form-item label="对话密码">
<n-input v-model:value="talkPassword" type="password" show-password-on="click" placeholder="留空表示清除" /> <n-input v-model:value="talkPassword" type="password" show-password-on="click" placeholder="4–64 字符,留空表示清除" />
</n-form-item> </n-form-item>
</n-form> </n-form>
<template #footer> <template #footer>
@@ -478,14 +654,38 @@ function onPageSizeChange(s: number) {
<n-modal v-if="importResult" :show="true" preset="card" title="导入结果" style="width: 560px" @update:show="(v) => !v && (importResult = null)"> <n-modal v-if="importResult" :show="true" preset="card" title="导入结果" style="width: 560px" @update:show="(v) => !v && (importResult = null)">
<n-scrollbar style="max-height: 280px"> <n-scrollbar style="max-height: 280px">
<n-text depth="3">以下密码仅本次可见,请下载保存。</n-text> <n-text depth="3">以下密码仅本次可见,请下载保存。</n-text>
<pre style="font-size: 12px; white-space: pre-wrap">{{ importResult.map((i) => `${i.id}\t${i.name}\t${i.login_password}`).join("\n") }}</pre> <pre style="font-size: 12px; white-space: pre-wrap">{{ importResult.map((i) => `${i.id}\t${i.name}\t${i.login_password ?? ""}`).join("\n") }}</pre>
</n-scrollbar> </n-scrollbar>
<template #footer> <template #footer>
<n-space justify="end"> <n-space justify="end">
<n-button @click="downloadImportCsv">下载</n-button>
<n-button @click="importResult = null">关闭</n-button> <n-button @click="importResult = null">关闭</n-button>
</n-space> </n-space>
</template> </template>
</n-modal> </n-modal>
<n-modal
:show="!!importErrors"
preset="card"
title="导入校验失败"
style="width: 560px"
@update:show="(v) => !v && (importErrors = null)"
>
<n-data-table
:columns="[
{ title: '行号', key: 'line', width: 80 },
{ title: '原因', key: 'reason' },
]"
:data="importErrors || []"
:max-height="280"
:row-key="(r: { line: number; reason: string }) => String(r.line) + r.reason"
data-testid="import-errors"
/>
<template #footer>
<n-space justify="end">
<n-button @click="importErrors = null">关闭</n-button>
</n-space>
</template>
</n-modal>
</div> </div>
</template> </template>
@@ -498,8 +698,17 @@ function onPageSizeChange(s: number) {
} }
.page-body { .page-body {
flex: 1; flex: 1;
overflow: auto; min-height: 0;
display: flex;
flex-direction: column;
overflow: hidden;
padding: 16px; padding: 16px;
}
.page-filters {
flex: 0 0 auto;
}
.page-table {
flex: 1;
min-height: 0; min-height: 0;
} }
</style> </style>
+77
View File
@@ -0,0 +1,77 @@
import { config, mount, flushPromises } from "@vue/test-utils";
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
NConfigProvider,
NDialogProvider,
NMessageProvider,
dateZhCN,
zhCN,
} from "naive-ui";
import { defineComponent, h } from "vue";
import GroupsView from "./GroupsView.vue";
import { mockApi } from "@/api/mock";
import * as admin from "@/api/admin";
vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
config.global.stubs = { teleport: true };
function wrap(Comp: object) {
return defineComponent({
setup() {
return () =>
h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, {
default: () =>
h(NMessageProvider, null, {
default: () =>
h(NDialogProvider, null, {
default: () => h(Comp),
}),
}),
});
},
});
}
describe("GroupsView", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
vi.restoreAllMocks();
});
it("搜索占位为群名称,无编号输入,详情可翻到第 2 页", async () => {
const getSpy = vi.spyOn(admin, "getGroup");
const w = mount(wrap(GroupsView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
expect(w.find('[data-testid="group-query"]').attributes("placeholder") || w.html()).toBeTruthy();
expect(w.find('[placeholder="群名称"]').exists() || w.html().includes("群名称")).toBe(true);
await w.find('[data-testid="group-create-open"]').trigger("click");
await flushPromises();
expect(w.find('[data-testid="group-id"]').exists()).toBe(false);
expect(w.text()).toContain("创建后自动分配");
const detailBtn = w.findAll("button").find((b) => b.text().trim() === "详情");
await detailBtn!.trigger("click");
await flushPromises();
expect(getSpy).toHaveBeenCalled();
getSpy.mockClear();
getSpy.mockImplementation((id, cursor) => mockApi.getGroup(id, cursor, 1));
await detailBtn!.trigger("click");
await flushPromises();
const page2 = w.findAll("li,button").find((el) => el.text().trim() === "2");
if (page2) {
await page2.trigger("click");
await flushPromises();
expect(getSpy.mock.calls.some((c) => String(c[1] || "").length > 0)).toBe(true);
}
w.unmount();
});
});
+106 -19
View File
@@ -10,9 +10,12 @@ import {
NModal, NModal,
NScrollbar, NScrollbar,
NSpace, NSpace,
NText,
useDialog, useDialog,
} from "naive-ui"; } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import HelpTip from "@/components/HelpTip.vue";
import { import {
addGroupMembers, addGroupMembers,
createGroup, createGroup,
@@ -31,6 +34,7 @@ import { message } from "@/utils/notify";
const dialog = useDialog(); const dialog = useDialog();
const loading = ref(false); const loading = ref(false);
const loadError = ref("");
const rows = ref<GroupSummary[]>([]); const rows = ref<GroupSummary[]>([]);
const total = ref(0); const total = ref(0);
const page = ref(1); const page = ref(1);
@@ -38,21 +42,27 @@ const pageSize = ref(50);
const query = ref(""); const query = ref("");
const createOpen = ref(false); const createOpen = ref(false);
const createForm = reactive({ id: "", name: "", owner_id: "", member_ids: "" }); const createForm = reactive({ name: "", owner_id: "", member_ids: "" });
const createLoading = ref(false); const createLoading = ref(false);
const detailOpen = ref(false); const detailOpen = ref(false);
const detailId = ref("");
const detail = ref<GroupDetail | null>(null); const detail = ref<GroupDetail | null>(null);
const memberPage = ref(1);
const memberPageSize = ref(50);
const renameName = ref(""); const renameName = ref("");
const addMembers = ref(""); const addMembers = ref("");
async function load() { async function load() {
loading.value = true; loading.value = true;
loadError.value = "";
try { try {
const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : ""; const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : "";
const res = await listGroups(query.value.trim() || undefined, cursor, pageSize.value); const res = await listGroups(query.value.trim() || undefined, cursor, pageSize.value);
rows.value = res.items; rows.value = res.items;
total.value = res.total; total.value = res.total;
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
@@ -132,13 +142,38 @@ const memberColumns = computed<DataTableColumns<GroupMember>>(() => [
}, },
]); ]);
async function openDetail(id: string) { async function loadDetail(id: string) {
detail.value = await getGroup(id); const cursor = memberPage.value > 1 ? String((memberPage.value - 1) * memberPageSize.value) : "";
detail.value = await getGroup(id, cursor, memberPageSize.value);
renameName.value = detail.value.name; renameName.value = detail.value.name;
}
async function openDetail(id: string) {
detailId.value = id;
memberPage.value = 1;
addMembers.value = ""; addMembers.value = "";
await loadDetail(id);
detailOpen.value = true; detailOpen.value = true;
} }
function failLabel(code: string): string {
switch (code) {
case "not_found":
case "invalid_target":
return "端不存在";
case "endpoint_disabled":
return "端已停用";
case "group_full":
return "群已满";
default:
return code;
}
}
function formatFailed(failed: { id: string; code: string }[]): string {
return failed.map((f) => `${f.id}(${failLabel(f.code)})`).join("、");
}
async function submitCreate() { async function submitCreate() {
if (!createForm.name.trim() || !createForm.owner_id.trim()) { if (!createForm.name.trim() || !createForm.owner_id.trim()) {
message.error("名称与群主必填"); message.error("名称与群主必填");
@@ -151,14 +186,13 @@ async function submitCreate() {
.map((s) => s.trim()) .map((s) => s.trim())
.filter(Boolean); .filter(Boolean);
const res = await createGroup({ const res = await createGroup({
id: createForm.id.trim() || undefined,
name: createForm.name.trim(), name: createForm.name.trim(),
owner_id: createForm.owner_id.trim(), owner_id: createForm.owner_id.trim(),
member_ids, member_ids,
}); });
createOpen.value = false; createOpen.value = false;
if (res.failed.length) { if (res.failed.length) {
message.warning(`已创建,部分成员失败:${res.failed.map((f) => f.id).join(", ")}`); message.warning(`已创建,部分成员失败:${formatFailed(res.failed)}`);
} else { } else {
message.success("已创建"); message.success("已创建");
} }
@@ -186,7 +220,7 @@ async function onRename() {
if (!detail.value) return; if (!detail.value) return;
await renameGroup(detail.value.id, renameName.value.trim()); await renameGroup(detail.value.id, renameName.value.trim());
message.success("已改名"); message.success("已改名");
await openDetail(detail.value.id); await loadDetail(detail.value.id);
await load(); await load();
} }
@@ -199,28 +233,44 @@ async function onAddMembers() {
if (!ids.length) return; if (!ids.length) return;
const res = await addGroupMembers(detail.value.id, ids); const res = await addGroupMembers(detail.value.id, ids);
if (res.failed.length) { if (res.failed.length) {
message.warning(`部分失败:${res.failed.map((f) => f.id).join(", ")}`); message.warning(`部分失败:${formatFailed(res.failed)}`);
} else { } else {
message.success("已加人"); message.success("已加人");
} }
await openDetail(detail.value.id); await loadDetail(detail.value.id);
await load(); await load();
} }
async function onRemove(endpointId: string) { async function onRemove(endpointId: string) {
if (!detail.value) return; if (!detail.value) return;
await removeGroupMember(detail.value.id, endpointId); dialog.warning({
title: "移除成员",
content: `确定将 ${endpointId} 移出本群?`,
positiveText: "移除",
negativeText: "取消",
onPositiveClick: async () => {
await removeGroupMember(detail.value!.id, endpointId);
message.success("已移除"); message.success("已移除");
await openDetail(detail.value.id); await loadDetail(detail.value!.id);
await load(); await load();
},
});
} }
async function onTransfer(endpointId: string) { async function onTransfer(endpointId: string) {
if (!detail.value) return; if (!detail.value) return;
await transferGroup(detail.value.id, endpointId); dialog.warning({
title: "转让群主",
content: `确定将群主转让给 ${endpointId}?`,
positiveText: "转让",
negativeText: "取消",
onPositiveClick: async () => {
await transferGroup(detail.value!.id, endpointId);
message.success("已转让群主"); message.success("已转让群主");
await openDetail(detail.value.id); await loadDetail(detail.value!.id);
await load(); await load();
},
});
} }
</script> </script>
@@ -232,7 +282,7 @@ async function onTransfer(endpointId: string) {
type="primary" type="primary"
data-testid="group-create-open" data-testid="group-create-open"
@click=" @click="
Object.assign(createForm, { id: '', name: '', owner_id: '', member_ids: '' }); Object.assign(createForm, { name: '', owner_id: '', member_ids: '' });
createOpen = true; createOpen = true;
" "
> >
@@ -241,17 +291,22 @@ async function onTransfer(endpointId: string) {
</template> </template>
</PageHeader> </PageHeader>
<div class="page-body"> <div class="page-body">
<LoadFailed v-if="loadError" :description="loadError" @retry="load" />
<template v-else>
<div class="page-filters">
<n-space style="margin-bottom: 12px"> <n-space style="margin-bottom: 12px">
<n-input v-model:value="query" placeholder="名称/编号" style="width: 200px" @keyup.enter="page = 1; load()" /> <n-input v-model:value="query" placeholder="群名称" style="width: 200px" data-testid="group-query" @keyup.enter="page = 1; load()" />
<n-button @click="page = 1; load()">查询</n-button> <n-button @click="page = 1; load()">查询</n-button>
</n-space> </n-space>
</div>
<div class="page-table">
<n-data-table <n-data-table
remote remote
flex-height
:loading="loading" :loading="loading"
:columns="columns" :columns="columns"
:data="rows" :data="rows"
:row-key="(r: GroupSummary) => r.id" :row-key="(r: GroupSummary) => r.id"
:max-height="480"
:pagination="{ :pagination="{
page, page,
pageSize, pageSize,
@@ -263,12 +318,18 @@ async function onTransfer(endpointId: string) {
}" }"
/> />
</div> </div>
</template>
</div>
<n-modal v-model:show="createOpen" preset="card" title="新建群" style="width: 480px"> <n-modal v-model:show="createOpen" preset="card" title="新建群" style="width: 480px">
<n-scrollbar style="max-height: 320px"> <n-scrollbar style="max-height: 320px">
<n-form label-placement="left" label-width="90" :show-feedback="false"> <n-form label-placement="left" label-width="90" :show-feedback="false">
<n-form-item label="编号"> <n-form-item>
<n-input v-model:value="createForm.id" placeholder="留空则生成" data-testid="group-id" /> <template #label>
编号
<HelpTip>由服务器生成,后台不能自定。</HelpTip>
</template>
<n-text depth="3">创建后自动分配</n-text>
</n-form-item> </n-form-item>
<n-form-item label="名称"> <n-form-item label="名称">
<n-input v-model:value="createForm.name" data-testid="group-name" /> <n-input v-model:value="createForm.name" data-testid="group-name" />
@@ -315,7 +376,23 @@ async function onTransfer(endpointId: string) {
</n-space> </n-space>
</n-form-item> </n-form-item>
</n-form> </n-form>
<n-data-table :columns="memberColumns" :data="detail.members" :row-key="(r: GroupMember) => r.id" size="small" /> <n-data-table
remote
:columns="memberColumns"
:data="detail.members"
:row-key="(r: GroupMember) => r.id"
size="small"
:max-height="280"
:pagination="{
page: memberPage,
pageSize: memberPageSize,
itemCount: detail.member_total ?? detail.members.length,
onChange: (p: number) => {
memberPage = p;
loadDetail(detailId);
},
}"
/>
</n-scrollbar> </n-scrollbar>
<template #footer> <template #footer>
<n-space justify="end"> <n-space justify="end">
@@ -335,7 +412,17 @@ async function onTransfer(endpointId: string) {
} }
.page-body { .page-body {
flex: 1; flex: 1;
overflow: auto; min-height: 0;
display: flex;
flex-direction: column;
overflow: hidden;
padding: 16px; padding: 16px;
} }
.page-filters {
flex: 0 0 auto;
}
.page-table {
flex: 1;
min-height: 0;
}
</style> </style>
+79
View File
@@ -0,0 +1,79 @@
import { config, mount, flushPromises } from "@vue/test-utils";
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
NConfigProvider,
NDialogProvider,
NMessageProvider,
dateZhCN,
zhCN,
} from "naive-ui";
import { defineComponent, h, nextTick } from "vue";
import MessagesView from "./MessagesView.vue";
import { mockApi } from "@/api/mock";
import * as admin from "@/api/admin";
vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
config.global.stubs = { teleport: true };
function wrap(Comp: object) {
return defineComponent({
setup() {
return () =>
h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, {
default: () =>
h(NMessageProvider, null, {
default: () =>
h(NDialogProvider, null, {
default: () => h(Comp),
}),
}),
});
},
});
}
describe("MessagesView", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
vi.restoreAllMocks();
});
it("查询时传递时间范围,统计含拒绝,详情有原因列并可翻页", async () => {
const listSpy = vi.spyOn(admin, "listMessages");
const getSpy = vi.spyOn(admin, "getMessage").mockImplementation((seq, cursor) => mockApi.getMessage(seq, cursor, 1));
const w = mount(wrap(MessagesView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
expect(w.text()).toMatch(/已拒绝/);
const picker = w.findComponent({ name: "DatePicker" });
expect(picker.exists()).toBe(true);
picker.vm.$emit("update:value", [1_700_000_000_000, 1_800_000_000_000]);
await nextTick();
const queryBtn = w.findAll("button").find((b) => b.text().trim() === "查询");
await queryBtn!.trigger("click");
await flushPromises();
const last = listSpy.mock.calls.at(-1)?.[0] as { from_ms?: number; to_ms?: number };
expect(last.from_ms).toBe(1_700_000_000_000);
expect(last.to_ms).toBe(1_800_000_000_000);
const detailBtn = w.findAll("button").find((b) => b.text().trim() === "详情");
await detailBtn!.trigger("click");
await flushPromises();
expect(w.text()).toContain("原因");
expect(w.find('[data-testid="delivery-more"]').exists()).toBe(true);
getSpy.mockClear();
await w.find('[data-testid="delivery-more"]').trigger("click");
await flushPromises();
expect(getSpy).toHaveBeenCalled();
expect(getSpy.mock.calls[0][1]).toBeTruthy();
w.unmount();
});
});
+86 -15
View File
@@ -4,20 +4,23 @@ import type { DataTableColumns } from "naive-ui";
import { import {
NButton, NButton,
NDataTable, NDataTable,
NDatePicker,
NDescriptions, NDescriptions,
NDescriptionsItem, NDescriptionsItem,
NInput, NInput,
NModal, NModal,
NScrollbar,
NSelect, NSelect,
NSpace, NSpace,
} from "naive-ui"; } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import HelpTip from "@/components/HelpTip.vue"; import HelpTip from "@/components/HelpTip.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import { getMessage, listMessages, type MessageDelivery, type MessageDetail, type MessageSummary } from "@/api/admin"; import { getMessage, listMessages, type MessageDelivery, type MessageDetail, type MessageSummary } from "@/api/admin";
import { formatLocalMs } from "@/utils/time"; import { formatLocalMs } from "@/utils/time";
import { destKindLabel, deliveryStateLabel, messageStateLabel, reasonLabel, zhLabel } from "@/utils/labels";
const loading = ref(false); const loading = ref(false);
const loadError = ref("");
const rows = ref<MessageSummary[]>([]); const rows = ref<MessageSummary[]>([]);
const total = ref(0); const total = ref(0);
const page = ref(1); const page = ref(1);
@@ -28,9 +31,12 @@ const filters = reactive({
group_id: "", group_id: "",
state: "" as "" | "scheduled" | "dispatched" | "completed", state: "" as "" | "scheduled" | "dispatched" | "completed",
}); });
const timeRange = ref<[number, number] | null>(null);
const detailOpen = ref(false); const detailOpen = ref(false);
const detail = ref<MessageDetail | null>(null); const detail = ref<MessageDetail | null>(null);
const detailSeq = ref(0);
const deliveryLoading = ref(false);
const stateOptions = [ const stateOptions = [
{ label: "全部状态", value: "" }, { label: "全部状态", value: "" },
@@ -41,6 +47,7 @@ const stateOptions = [
async function load() { async function load() {
loading.value = true; loading.value = true;
loadError.value = "";
try { try {
const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : ""; const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : "";
const res = await listMessages({ const res = await listMessages({
@@ -50,9 +57,13 @@ async function load() {
endpoint_id: filters.endpoint_id.trim() || undefined, endpoint_id: filters.endpoint_id.trim() || undefined,
group_id: filters.group_id.trim() || undefined, group_id: filters.group_id.trim() || undefined,
state: filters.state || undefined, state: filters.state || undefined,
from_ms: timeRange.value?.[0],
to_ms: timeRange.value?.[1],
}); });
rows.value = res.items; rows.value = res.items;
total.value = res.total; total.value = res.total;
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
@@ -70,16 +81,23 @@ const columns = computed<DataTableColumns<MessageSummary>>(() => [
title: "目标", title: "目标",
key: "dest", key: "dest",
width: 140, width: 140,
render: (r) => `${r.dest_kind}:${r.dest_id}`, render: (r) => `${zhLabel(destKindLabel, r.dest_kind)}:${r.dest_id}`,
},
{ title: "状态", key: "state", width: 100, render: (r) => h("span", { title: r.state }, zhLabel(messageStateLabel, r.state)) },
{
title: "原因",
key: "reason",
width: 140,
ellipsis: { tooltip: true },
render: (r) => h("span", { title: r.reason || "" }, r.reason ? zhLabel(reasonLabel, r.reason) : "—"),
}, },
{ title: "状态", key: "state", width: 100 },
{ {
title: "投递统计", title: "投递统计",
key: "delivery_counts", key: "delivery_counts",
width: 200, width: 280,
render: (r) => { render: (r) => {
const c = r.delivery_counts; const c = r.delivery_counts;
return `待${c.pending}/收${c.accepted}/撤${c.recalled}/过${c.expired}`; return `待投递${c.pending}/已收下${c.accepted}/已撤回${c.recalled}/已过期${c.expired}/已丢弃${c.dropped}/已拒绝${c.rejected}`;
}, },
}, },
{ {
@@ -104,7 +122,12 @@ const columns = computed<DataTableColumns<MessageSummary>>(() => [
const deliveryColumns = computed<DataTableColumns<MessageDelivery>>(() => [ const deliveryColumns = computed<DataTableColumns<MessageDelivery>>(() => [
{ title: "接收端", key: "endpoint_id" }, { title: "接收端", key: "endpoint_id" },
{ title: "状态", key: "state" }, { title: "状态", key: "state", render: (r) => h("span", { title: r.state }, zhLabel(deliveryStateLabel, r.state)) },
{
title: "原因",
key: "reason",
render: (r) => h("span", { title: r.reason || "" }, r.reason ? zhLabel(reasonLabel, r.reason) : "—"),
},
{ title: "次数", key: "attempts", width: 60 }, { title: "次数", key: "attempts", width: 60 },
{ {
title: "推送时间", title: "推送时间",
@@ -119,20 +142,44 @@ const deliveryColumns = computed<DataTableColumns<MessageDelivery>>(() => [
]); ]);
async function openDetail(seq: number) { async function openDetail(seq: number) {
detail.value = await getMessage(seq); detailSeq.value = seq;
detail.value = await getMessage(seq, "", 200);
detailOpen.value = true; detailOpen.value = true;
} }
async function loadMoreDeliveries() {
if (!detail.value?.next_cursor) return;
deliveryLoading.value = true;
try {
const more = await getMessage(detailSeq.value, detail.value.next_cursor, 200);
detail.value = {
...more,
deliveries: [...detail.value.deliveries, ...more.deliveries],
};
} finally {
deliveryLoading.value = false;
}
}
</script> </script>
<template> <template>
<div class="page"> <div class="page">
<PageHeader title="投递记录" help="只显示投递痕迹,不显示正文;接口响应也不含 body 字段。" /> <PageHeader title="投递记录" help="只显示投递痕迹,不显示正文;接口响应也不含 body 字段。" />
<div class="page-body"> <div class="page-body">
<LoadFailed v-if="loadError" :description="loadError" @retry="load" />
<template v-else>
<div class="page-filters">
<n-space style="margin-bottom: 12px" align="center"> <n-space style="margin-bottom: 12px" align="center">
<n-input v-model:value="filters.sender_id" placeholder="发送方" style="width: 120px" /> <n-input v-model:value="filters.sender_id" placeholder="发送方" style="width: 120px" />
<n-input v-model:value="filters.endpoint_id" placeholder="接收端" style="width: 120px" /> <n-input v-model:value="filters.endpoint_id" placeholder="接收端" style="width: 120px" />
<n-input v-model:value="filters.group_id" placeholder="群编号" style="width: 120px" /> <n-input v-model:value="filters.group_id" placeholder="群编号" style="width: 120px" />
<n-select v-model:value="filters.state" :options="stateOptions" style="width: 120px" /> <n-select v-model:value="filters.state" :options="stateOptions" style="width: 120px" />
<n-date-picker
v-model:value="timeRange"
type="datetimerange"
clearable
data-testid="msg-time-range"
/>
<n-button <n-button
@click=" @click="
page = 1; page = 1;
@@ -146,14 +193,16 @@ async function openDetail(seq: number) {
<HelpTip>页面与网络响应均不得出现消息正文。</HelpTip> <HelpTip>页面与网络响应均不得出现消息正文。</HelpTip>
</span> </span>
</n-space> </n-space>
</div>
<div class="page-table">
<n-data-table <n-data-table
remote remote
flex-height
:loading="loading" :loading="loading"
:columns="columns" :columns="columns"
:data="rows" :data="rows"
:row-key="(r: MessageSummary) => r.seq" :row-key="(r: MessageSummary) => r.seq"
:max-height="480" :scroll-x="1400"
:scroll-x="1100"
:pagination="{ :pagination="{
page, page,
pageSize, pageSize,
@@ -165,21 +214,33 @@ async function openDetail(seq: number) {
}" }"
/> />
</div> </div>
</template>
</div>
<n-modal v-model:show="detailOpen" preset="card" title="投递详情" style="width: 720px"> <n-modal v-model:show="detailOpen" preset="card" title="投递详情" style="width: 720px">
<n-scrollbar v-if="detail" style="max-height: 420px"> <template v-if="detail">
<n-descriptions label-placement="left" :column="2" size="small" bordered style="margin-bottom: 12px"> <n-descriptions label-placement="left" :column="2" size="small" bordered style="margin-bottom: 12px">
<n-descriptions-item label="序号">{{ detail.seq }}</n-descriptions-item> <n-descriptions-item label="序号">{{ detail.seq }}</n-descriptions-item>
<n-descriptions-item label="消息号">{{ detail.id }}</n-descriptions-item> <n-descriptions-item label="消息号">{{ detail.id }}</n-descriptions-item>
<n-descriptions-item label="发送方">{{ detail.sender_id }}</n-descriptions-item> <n-descriptions-item label="发送方">{{ detail.sender_id }}</n-descriptions-item>
<n-descriptions-item label="目标">{{ detail.dest_kind }}:{{ detail.dest_id }}</n-descriptions-item> <n-descriptions-item label="目标">{{ destKindLabel[detail.dest_kind] || detail.dest_kind }}:{{ detail.dest_id }}</n-descriptions-item>
<n-descriptions-item label="状态">{{ detail.state }}</n-descriptions-item> <n-descriptions-item label="状态"><span :title="detail.state">{{ zhLabel(messageStateLabel, detail.state) }}</span></n-descriptions-item>
<n-descriptions-item label="原因"><span :title="detail.reason">{{ detail.reason ? zhLabel(reasonLabel, detail.reason) : "—" }}</span></n-descriptions-item>
<n-descriptions-item label="类型">{{ detail.content_type }}</n-descriptions-item> <n-descriptions-item label="类型">{{ detail.content_type }}</n-descriptions-item>
<n-descriptions-item label="创建">{{ formatLocalMs(detail.created_at_ms) }}</n-descriptions-item> <n-descriptions-item label="创建">{{ formatLocalMs(detail.created_at_ms) }}</n-descriptions-item>
<n-descriptions-item label="发送">{{ formatLocalMs(detail.send_at_ms) }}</n-descriptions-item> <n-descriptions-item label="发送">{{ formatLocalMs(detail.send_at_ms) }}</n-descriptions-item>
</n-descriptions> </n-descriptions>
<n-data-table :columns="deliveryColumns" :data="detail.deliveries" size="small" :row-key="(r: MessageDelivery) => r.endpoint_id" /> <n-data-table :columns="deliveryColumns" :data="detail.deliveries" size="small" :row-key="(r: MessageDelivery) => r.endpoint_id" :max-height="280" />
</n-scrollbar> <n-button
v-if="detail.next_cursor"
:loading="deliveryLoading"
data-testid="delivery-more"
style="margin-top: 8px"
@click="loadMoreDeliveries"
>
加载更多
</n-button>
</template>
<template #footer> <template #footer>
<n-space justify="end"> <n-space justify="end">
<n-button @click="detailOpen = false">关闭</n-button> <n-button @click="detailOpen = false">关闭</n-button>
@@ -198,7 +259,17 @@ async function openDetail(seq: number) {
} }
.page-body { .page-body {
flex: 1; flex: 1;
overflow: auto; min-height: 0;
display: flex;
flex-direction: column;
overflow: hidden;
padding: 16px; padding: 16px;
} }
.page-filters {
flex: 0 0 auto;
}
.page-table {
flex: 1;
min-height: 0;
}
</style> </style>
+85
View File
@@ -0,0 +1,85 @@
import { config, mount, flushPromises } from "@vue/test-utils";
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
NConfigProvider,
NDialogProvider,
NMessageProvider,
dateZhCN,
zhCN,
} from "naive-ui";
import { defineComponent, h } from "vue";
import OverviewView from "./OverviewView.vue";
import * as admin from "@/api/admin";
import { mockApi } from "@/api/mock";
vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
config.global.stubs = { teleport: true };
function wrap(Comp: object) {
return defineComponent({
setup() {
return () =>
h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, {
default: () =>
h(NMessageProvider, null, {
default: () =>
h(NDialogProvider, null, {
default: () => h(Comp),
}),
}),
});
},
});
}
describe("OverviewView", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
vi.restoreAllMocks();
});
it("失败时显示重试,点后重新加载", async () => {
const spy = vi.spyOn(admin, "fetchOverview");
spy.mockRejectedValueOnce(new Error("boom")).mockResolvedValueOnce({
version: "0.1.0",
endpoints_total: 3,
endpoints_online: 1,
endpoints_disabled: 1,
endpoints_self: 1,
groups_total: 1,
messages_pending: 0,
messages_scheduled: 0,
uptime_ms: 1000,
});
const w = mount(wrap(OverviewView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
expect(w.text()).toContain("加载失败");
expect(w.find('[data-testid="load-retry"]').exists()).toBe(true);
await w.find('[data-testid="load-retry"]').trigger("click");
await flushPromises();
expect(w.text()).toContain("0.1.0");
expect(w.text()).toContain("其中自助注册");
expect(w.text()).not.toContain("加载失败");
w.unmount();
});
it("成功加载后显示自助注册数", async () => {
const w = mount(wrap(OverviewView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
expect(w.text()).toContain("其中自助注册");
w.unmount();
});
});
+16 -3
View File
@@ -3,18 +3,29 @@ import { onMounted, ref } from "vue";
import { NDescriptions, NDescriptionsItem, NSpin } from "naive-ui"; import { NDescriptions, NDescriptionsItem, NSpin } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import HelpTip from "@/components/HelpTip.vue"; import HelpTip from "@/components/HelpTip.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import { fetchOverview, type Overview } from "@/api/admin"; import { fetchOverview, type Overview } from "@/api/admin";
import { formatLocalMs } from "@/utils/time"; import { formatLocalMs } from "@/utils/time";
import { formatUptime } from "@/utils/labels";
const loading = ref(true); const loading = ref(true);
const loadError = ref("");
const data = ref<Overview | null>(null); const data = ref<Overview | null>(null);
onMounted(async () => { async function load() {
loading.value = true;
loadError.value = "";
try { try {
data.value = await fetchOverview(); data.value = await fetchOverview();
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
}
onMounted(() => {
void load();
}); });
</script> </script>
@@ -23,16 +34,18 @@ onMounted(async () => {
<PageHeader title="概览" help="汇总当前实例的端、群与未完成投递数量,不含正文与编号明细。" /> <PageHeader title="概览" help="汇总当前实例的端、群与未完成投递数量,不含正文与编号明细。" />
<div class="page-body"> <div class="page-body">
<n-spin :show="loading"> <n-spin :show="loading">
<n-descriptions v-if="data" label-placement="left" :column="2" bordered size="small"> <LoadFailed v-if="loadError && !data" :description="loadError" @retry="load" />
<n-descriptions v-else-if="data" label-placement="left" :column="2" bordered size="small">
<n-descriptions-item label="版本">{{ data.version }}</n-descriptions-item> <n-descriptions-item label="版本">{{ data.version }}</n-descriptions-item>
<n-descriptions-item> <n-descriptions-item>
<template #label> <template #label>
运行时长 运行时长
<HelpTip>自进程启动起的毫秒数,按本地时区换算展示起点附近时间差。</HelpTip> <HelpTip>自进程启动起的毫秒数,按本地时区换算展示起点附近时间差。</HelpTip>
</template> </template>
{{ Math.floor(data.uptime_ms / 1000) }} 秒(约 {{ formatLocalMs(Date.now() - data.uptime_ms) }} 起) {{ formatUptime(data.uptime_ms) }}(约 {{ formatLocalMs(Date.now() - data.uptime_ms) }} 起)
</n-descriptions-item> </n-descriptions-item>
<n-descriptions-item label="端总数">{{ data.endpoints_total }}</n-descriptions-item> <n-descriptions-item label="端总数">{{ data.endpoints_total }}</n-descriptions-item>
<n-descriptions-item label="其中自助注册">{{ data.endpoints_self }}</n-descriptions-item>
<n-descriptions-item label="在线端">{{ data.endpoints_online }}</n-descriptions-item> <n-descriptions-item label="在线端">{{ data.endpoints_online }}</n-descriptions-item>
<n-descriptions-item label="已停用">{{ data.endpoints_disabled }}</n-descriptions-item> <n-descriptions-item label="已停用">{{ data.endpoints_disabled }}</n-descriptions-item>
<n-descriptions-item label="群数量">{{ data.groups_total }}</n-descriptions-item> <n-descriptions-item label="群数量">{{ data.groups_total }}</n-descriptions-item>
+53
View File
@@ -0,0 +1,53 @@
import { config, mount, flushPromises } from "@vue/test-utils";
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
NConfigProvider,
NDialogProvider,
NMessageProvider,
dateZhCN,
zhCN,
} from "naive-ui";
import { defineComponent, h } from "vue";
import RegistrationView from "./RegistrationView.vue";
import { mockApi } from "@/api/mock";
vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
config.global.stubs = { teleport: true };
function wrap(Comp: object) {
return defineComponent({
setup() {
return () =>
h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, {
default: () =>
h(NMessageProvider, null, {
default: () =>
h(NDialogProvider, null, {
default: () => h(Comp),
}),
}),
});
},
});
}
describe("RegistrationView", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
});
it("已有安全码时不显示缺码提示", async () => {
const pinia = createPinia();
const w = mount(wrap(RegistrationView), {
global: { plugins: [pinia] },
attachTo: document.body,
});
await flushPromises();
expect(w.get('[data-testid="reg-enabled"]').exists()).toBe(true);
expect(w.find('[data-testid="reg-code-missing"]').exists()).toBe(false);
w.unmount();
});
});
+18 -3
View File
@@ -1,6 +1,6 @@
<script setup lang="ts"> <script setup lang="ts">
import { onMounted, ref } from "vue"; import { computed, onMounted, ref } from "vue";
import { NButton, NForm, NFormItem, NInput, NSpace, NSpin, NSwitch, NText } from "naive-ui"; import { NAlert, NButton, NForm, NFormItem, NInput, NSpace, NSpin, NSwitch, NText } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import HelpTip from "@/components/HelpTip.vue"; import HelpTip from "@/components/HelpTip.vue";
import { getRegistration, updateRegistration, type RegistrationSettings } from "@/api/admin"; import { getRegistration, updateRegistration, type RegistrationSettings } from "@/api/admin";
@@ -12,6 +12,9 @@ const saving = ref(false);
const data = ref<RegistrationSettings | null>(null); const data = ref<RegistrationSettings | null>(null);
const codeDraft = ref(""); const codeDraft = ref("");
const hasSavedCode = computed(() => (data.value?.code ?? "").length >= 8);
const switchDisabled = computed(() => !hasSavedCode.value && !data.value?.enabled);
async function load() { async function load() {
loading.value = true; loading.value = true;
try { try {
@@ -28,6 +31,10 @@ onMounted(() => {
async function onToggle(enabled: boolean) { async function onToggle(enabled: boolean) {
if (!data.value) return; if (!data.value) return;
if (enabled && !hasSavedCode.value) {
message.error("请先保存或生成安全码");
return;
}
saving.value = true; saving.value = true;
try { try {
data.value = await updateRegistration({ enabled }); data.value = await updateRegistration({ enabled });
@@ -66,6 +73,13 @@ async function generateCode() {
<div class="page-body"> <div class="page-body">
<n-spin :show="loading"> <n-spin :show="loading">
<n-form v-if="data" label-placement="left" label-width="100" :show-feedback="false" style="max-width: 560px"> <n-form v-if="data" label-placement="left" label-width="100" :show-feedback="false" style="max-width: 560px">
<n-alert
v-if="!hasSavedCode"
type="warning"
title="请先保存或生成安全码"
data-testid="reg-code-missing"
style="margin-bottom: 12px"
/>
<n-form-item> <n-form-item>
<template #label> <template #label>
开放注册 开放注册
@@ -74,6 +88,7 @@ async function generateCode() {
<n-switch <n-switch
:value="data.enabled" :value="data.enabled"
:loading="saving" :loading="saving"
:disabled="switchDisabled"
data-testid="reg-enabled" data-testid="reg-enabled"
@update:value="onToggle" @update:value="onToggle"
/> />
@@ -88,7 +103,7 @@ async function generateCode() {
<n-button type="primary" :loading="saving" data-testid="reg-save-code" @click="saveCode"> <n-button type="primary" :loading="saving" data-testid="reg-save-code" @click="saveCode">
保存安全码 保存安全码
</n-button> </n-button>
<n-button :loading="saving" @click="generateCode">生成安全码</n-button> <n-button :loading="saving" data-testid="reg-generate-code" @click="generateCode">生成安全码</n-button>
</n-space> </n-space>
</n-form> </n-form>
</n-spin> </n-spin>
+60
View File
@@ -0,0 +1,60 @@
import { config, mount, flushPromises } from "@vue/test-utils";
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import {
NConfigProvider,
NDialogProvider,
NMessageProvider,
dateZhCN,
zhCN,
} from "naive-ui";
import { defineComponent, h } from "vue";
import SettingsView from "./SettingsView.vue";
import { mockApi } from "@/api/mock";
vi.mock("@/api/admin", async () => import("@/api/admin-mock"));
config.global.stubs = { teleport: true };
function wrap(Comp: object) {
return defineComponent({
setup() {
return () =>
h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, {
default: () =>
h(NMessageProvider, null, {
default: () =>
h(NDialogProvider, null, {
default: () => h(Comp),
}),
}),
});
},
});
}
describe("SettingsView 改密表单", () => {
beforeEach(() => {
setActivePinia(createPinia());
mockApi._setSession();
});
it("新密码不足 12 个字符时就地显示错误", async () => {
const w = mount(wrap(SettingsView), {
global: { plugins: [createPinia()] },
attachTo: document.body,
});
await flushPromises();
const inputs = w.findAll("input");
expect(inputs.length).toBeGreaterThanOrEqual(3);
await inputs[0].setValue("old-password");
await inputs[1].setValue("short");
await inputs[2].setValue("short");
await w.findAll("button").find((b) => b.text().includes("修改密码"))!.trigger("click");
await flushPromises();
expect(w.text()).toContain("新密码至少 12 个字符");
w.unmount();
});
});
+97 -22
View File
@@ -1,6 +1,8 @@
<script setup lang="ts"> <script setup lang="ts">
import { onMounted, reactive, ref } from "vue"; import { onMounted, reactive, ref } from "vue";
import type { FormInst, FormRules } from "naive-ui";
import { import {
NAlert,
NButton, NButton,
NDescriptions, NDescriptions,
NDescriptionsItem, NDescriptionsItem,
@@ -15,29 +17,66 @@ import {
} from "naive-ui"; } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import HelpTip from "@/components/HelpTip.vue"; import HelpTip from "@/components/HelpTip.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import { changePassword, getSettings, type RuntimeSettings } from "@/api/admin"; import { changePassword, getSettings, type RuntimeSettings } from "@/api/admin";
import { ApiError } from "@/api/http";
import { message } from "@/utils/notify"; import { message } from "@/utils/notify";
import { limitsLabel } from "@/utils/labels";
const loading = ref(true); const loading = ref(true);
const loadError = ref("");
const settings = ref<RuntimeSettings | null>(null); const settings = ref<RuntimeSettings | null>(null);
const pwd = reactive({ old_password: "", new_password: "", confirm: "" }); const pwd = reactive({ old_password: "", new_password: "", confirm: "" });
const pwdLoading = ref(false); const pwdLoading = ref(false);
const pwdError = ref("");
const formRef = ref<FormInst | null>(null);
onMounted(async () => { const rules: FormRules = {
new_password: [
{
validator(_, value: string) {
if (Array.from(value || "").length < 12) {
return new Error("新密码至少 12 个字符");
}
return true;
},
trigger: ["blur", "input"],
},
],
confirm: [
{
validator() {
if (pwd.new_password !== pwd.confirm) {
return new Error("两次输入的新密码不一致");
}
return true;
},
trigger: ["blur", "input"],
},
],
};
async function load() {
loading.value = true;
loadError.value = "";
try { try {
settings.value = await getSettings(); settings.value = await getSettings();
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
}
onMounted(() => {
void load();
}); });
async function submitPassword() { async function submitPassword() {
if (pwd.new_password.length < 12) { pwdError.value = "";
message.error("新密码至少 12 位"); try {
return; await formRef.value?.validate();
} } catch {
if (pwd.new_password !== pwd.confirm) {
message.error("两次输入的新密码不一致");
return; return;
} }
pwdLoading.value = true; pwdLoading.value = true;
@@ -47,6 +86,8 @@ async function submitPassword() {
pwd.old_password = ""; pwd.old_password = "";
pwd.new_password = ""; pwd.new_password = "";
pwd.confirm = ""; pwd.confirm = "";
} catch (e) {
pwdError.value = e instanceof ApiError ? e.message : e instanceof Error ? e.message : "修改失败";
} finally { } finally {
pwdLoading.value = false; pwdLoading.value = false;
} }
@@ -59,18 +100,26 @@ async function submitPassword() {
<div class="page-body"> <div class="page-body">
<n-tabs type="line" animated> <n-tabs type="line" animated>
<n-tab-pane name="password" tab="管理员密码"> <n-tab-pane name="password" tab="管理员密码">
<n-form label-placement="left" label-width="100" :show-feedback="false" style="max-width: 480px"> <n-alert v-if="pwdError" type="error" style="margin-bottom: 12px; max-width: 480px" :title="pwdError" />
<n-form-item label="原密码"> <n-form
ref="formRef"
:model="pwd"
:rules="rules"
label-placement="left"
label-width="100"
style="max-width: 480px"
>
<n-form-item label="原密码" path="old_password">
<n-input v-model:value="pwd.old_password" type="password" show-password-on="click" /> <n-input v-model:value="pwd.old_password" type="password" show-password-on="click" />
</n-form-item> </n-form-item>
<n-form-item> <n-form-item path="new_password">
<template #label> <template #label>
新密码 新密码
<HelpTip>至少 12 位,且不能以 nst_ 开头(与端登录密码规则一致的安全要求)。</HelpTip> <HelpTip>按字符数计至少 12 位。后端不检查 nst_ 前缀(那是端登录密码规则)。</HelpTip>
</template> </template>
<n-input v-model:value="pwd.new_password" type="password" show-password-on="click" /> <n-input v-model:value="pwd.new_password" type="password" show-password-on="click" />
</n-form-item> </n-form-item>
<n-form-item label="确认新密码"> <n-form-item label="确认新密码" path="confirm">
<n-input v-model:value="pwd.confirm" type="password" show-password-on="click" /> <n-input v-model:value="pwd.confirm" type="password" show-password-on="click" />
</n-form-item> </n-form-item>
<n-space> <n-space>
@@ -80,19 +129,45 @@ async function submitPassword() {
</n-tab-pane> </n-tab-pane>
<n-tab-pane name="settings" tab="运行参数(只读)"> <n-tab-pane name="settings" tab="运行参数(只读)">
<n-spin :show="loading"> <n-spin :show="loading">
<template v-if="settings"> <LoadFailed v-if="loadError && !settings" :description="loadError" @retry="load" />
<template v-else-if="settings">
<n-descriptions label-placement="left" :column="2" bordered size="small" style="margin-bottom: 16px"> <n-descriptions label-placement="left" :column="2" bordered size="small" style="margin-bottom: 16px">
<n-descriptions-item label="listen">{{ settings.listen }}</n-descriptions-item> <n-descriptions-item>
<n-descriptions-item label="admin_listen">{{ settings.admin_listen || "(空,与端共用)" }}</n-descriptions-item> <template #label>端监听地址 <HelpTip>listen</HelpTip></template>
<n-descriptions-item label="会话闲置天数">{{ settings.session_idle_days }}</n-descriptions-item> {{ settings.listen }}
<n-descriptions-item label="记录保留天数">{{ settings.record_retention_days }}</n-descriptions-item> </n-descriptions-item>
<n-descriptions-item label="防重小时数">{{ settings.idempotency_hours }}</n-descriptions-item> <n-descriptions-item>
<n-descriptions-item label="回执保留天数">{{ settings.receipt_retention_days }}</n-descriptions-item> <template #label>后台监听地址 <HelpTip>admin_listen</HelpTip></template>
<n-descriptions-item label="SQLite synchronous">{{ settings.sqlite_synchronous }}</n-descriptions-item> {{ settings.admin_listen || "(空,与端共用)" }}
</n-descriptions-item>
<n-descriptions-item>
<template #label>会话闲置天数 <HelpTip>session_idle_days</HelpTip></template>
{{ settings.session_idle_days }}
</n-descriptions-item>
<n-descriptions-item>
<template #label>记录保留天数 <HelpTip>record_retention_days</HelpTip></template>
{{ settings.record_retention_days }}
</n-descriptions-item>
<n-descriptions-item>
<template #label>防重小时数 <HelpTip>idempotency_hours</HelpTip></template>
{{ settings.idempotency_hours }}
</n-descriptions-item>
<n-descriptions-item>
<template #label>回执保留天数 <HelpTip>receipt_retention_days</HelpTip></template>
{{ settings.receipt_retention_days }}
</n-descriptions-item>
<n-descriptions-item>
<template #label>SQLite 同步模式 <HelpTip>sqlite_synchronous</HelpTip></template>
{{ settings.sqlite_synchronous }}
</n-descriptions-item>
</n-descriptions> </n-descriptions>
<n-text strong style="display: block; margin-bottom: 8px">limits</n-text> <n-text strong style="display: block; margin-bottom: 8px">限额</n-text>
<n-descriptions label-placement="left" :column="2" bordered size="small"> <n-descriptions label-placement="left" :column="2" bordered size="small">
<n-descriptions-item v-for="(v, k) in settings.limits" :key="k" :label="String(k)"> <n-descriptions-item v-for="(v, k) in settings.limits" :key="k">
<template #label>
{{ limitsLabel[String(k)] || k }}
<HelpTip>{{ k }}</HelpTip>
</template>
{{ v }} {{ v }}
</n-descriptions-item> </n-descriptions-item>
</n-descriptions> </n-descriptions>
+24 -2
View File
@@ -12,6 +12,7 @@ import {
useDialog, useDialog,
} from "naive-ui"; } from "naive-ui";
import PageHeader from "@/components/PageHeader.vue"; import PageHeader from "@/components/PageHeader.vue";
import LoadFailed from "@/components/LoadFailed.vue";
import SecretOnceAlert from "@/components/SecretOnceAlert.vue"; import SecretOnceAlert from "@/components/SecretOnceAlert.vue";
import { createToken, deleteToken, listTokens, patchToken, type ApiToken } from "@/api/admin"; import { createToken, deleteToken, listTokens, patchToken, type ApiToken } from "@/api/admin";
import { formatLocalMs } from "@/utils/time"; import { formatLocalMs } from "@/utils/time";
@@ -19,6 +20,7 @@ import { message } from "@/utils/notify";
const dialog = useDialog(); const dialog = useDialog();
const loading = ref(false); const loading = ref(false);
const loadError = ref("");
const rows = ref<ApiToken[]>([]); const rows = ref<ApiToken[]>([]);
const createOpen = ref(false); const createOpen = ref(false);
const createName = ref(""); const createName = ref("");
@@ -27,9 +29,12 @@ const onceToken = ref<string | null>(null);
async function load() { async function load() {
loading.value = true; loading.value = true;
loadError.value = "";
try { try {
const res = await listTokens(); const res = await listTokens();
rows.value = res.items; rows.value = res.items;
} catch (e) {
loadError.value = e instanceof Error ? e.message : "加载失败";
} finally { } finally {
loading.value = false; loading.value = false;
} }
@@ -136,6 +141,9 @@ function onDelete(r: ApiToken) {
</template> </template>
</PageHeader> </PageHeader>
<div class="page-body"> <div class="page-body">
<LoadFailed v-if="loadError" :description="loadError" @retry="load" />
<template v-else>
<div class="page-filters">
<SecretOnceAlert <SecretOnceAlert
v-if="onceToken" v-if="onceToken"
title="令牌只显示一次" title="令牌只显示一次"
@@ -144,14 +152,18 @@ function onDelete(r: ApiToken) {
data-testid="token-once" data-testid="token-once"
@dismiss="onceToken = null" @dismiss="onceToken = null"
/> />
</div>
<div class="page-table">
<n-data-table <n-data-table
flex-height
:loading="loading" :loading="loading"
:columns="columns" :columns="columns"
:data="rows" :data="rows"
:row-key="(r: ApiToken) => r.id" :row-key="(r: ApiToken) => r.id"
:max-height="480"
/> />
</div> </div>
</template>
</div>
<n-modal v-model:show="createOpen" preset="card" title="创建 API 令牌" style="width: 420px"> <n-modal v-model:show="createOpen" preset="card" title="创建 API 令牌" style="width: 420px">
<n-form label-placement="left" label-width="60" :show-feedback="false"> <n-form label-placement="left" label-width="60" :show-feedback="false">
@@ -180,7 +192,17 @@ function onDelete(r: ApiToken) {
} }
.page-body { .page-body {
flex: 1; flex: 1;
overflow: auto; min-height: 0;
display: flex;
flex-direction: column;
overflow: hidden;
padding: 16px; padding: 16px;
} }
.page-filters {
flex: 0 0 auto;
}
.page-table {
flex: 1;
min-height: 0;
}
</style> </style>