merge: identity u01-u05
This commit is contained in:
@@ -1237,3 +1237,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`。
|
||||||
|
|||||||
@@ -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。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -141,6 +141,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 +260,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 +539,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")
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
|||||||
+121
-8
@@ -21,23 +21,34 @@ 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 是内存锁定计数器(重启清零)。
|
||||||
type MemoryLocks struct {
|
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 创建默认锁定计数器。
|
||||||
func NewLoginLocks() *MemoryLocks {
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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/);
|
||||||
|
|||||||
@@ -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();
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -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>
|
||||||
|
|||||||
Reference in New Issue
Block a user