diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index ddbc7e5..8e1c47e 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1237,3 +1237,30 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 5. **建议的正确方向** - 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。 - 覆盖 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`。 diff --git a/docs/api/admin-api.md b/docs/api/admin-api.md index 37abe91..877d68c 100644 --- a/docs/api/admin-api.md +++ b/docs/api/admin-api.md @@ -396,6 +396,8 @@ Authorization: Bearer nxm_... ``` - `generate: true` 时服务器生成 16 位安全码并忽略请求里的 `code`。 +- 允许一次提交 `{"enabled":true,"code":"..."}` 或 `{"enabled":true,"generate":true}`。 +- 最终状态为开启时,安全码必须是 8–64 字符;仅 `{"enabled":true}` 且库中没有合法安全码时返回 400,整次更新回滚。 - 成功 `data` 同 GET。 --- diff --git a/internal/admin/a3_test.go b/internal/admin/a3_test.go index 14f55df..953e3bb 100644 --- a/internal/admin/a3_test.go +++ b/internal/admin/a3_test.go @@ -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) { db, srv, client := setupA3(t) base := srv.URL diff --git a/internal/admin/registration.go b/internal/admin/registration.go index 583881c..c71415f 100644 --- a/internal/admin/registration.go +++ b/internal/admin/registration.go @@ -77,9 +77,7 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request) ); e != nil { return e } - return nil - } - if req.Code != nil { + } else if req.Code != nil { code := *req.Code n := utf8.RuneCountInString(code) if n < minRegistrationCodeLen || n > maxRegistrationCodeLen { @@ -93,6 +91,18 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request) 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 }) if err != nil { diff --git a/internal/app/identity/register.go b/internal/app/identity/register.go index 5f27b52..c64cbc7 100644 --- a/internal/app/identity/register.go +++ b/internal/app/identity/register.go @@ -25,6 +25,9 @@ const ( settingRegistrationEnabled = "registration_enabled" settingRegistrationCode = "registration_code" + minRegistrationCodeLen = 8 + maxRegistrationCodeLen = 64 + sourceSelf = "self" idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789" @@ -144,6 +147,11 @@ func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRe if !enabled { 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} 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) } + if strings.EqualFold(req.ID, "inline") { + return RegisterResult{}, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "id reserved") + } id := req.ID loginPassword := req.LoginPassword diff --git a/internal/app/identity/register_test.go b/internal/app/identity/register_test.go index 2d992eb..db51c04 100644 --- a/internal/app/identity/register_test.go +++ b/internal/app/identity/register_test.go @@ -141,6 +141,33 @@ func openTestEnv(t *testing.T) *testEnv { 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) { t.Helper() 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) { env := openTestEnv(t) 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) { env := openTestEnv(t) env.setRegistration(t, true, "nst-code-01") diff --git a/internal/app/presence/app.go b/internal/app/presence/app.go index f75966f..e5c1ddd 100644 --- a/internal/app/presence/app.go +++ b/internal/app/presence/app.go @@ -122,13 +122,15 @@ WHERE id > ? ORDER BY id ASC LIMIT ?`, cursor, limit+1) } else { - like := "%" + strings.ToLower(query) + "%" - prefix := strings.ToLower(query) + "%" + q := strings.ToLower(query) + esc := escapeLikePattern(q) + like := "%" + esc + "%" + prefix := esc + "%" rows, err = a.db.Read.QueryContext(ctx, ` SELECT id, name, online_since, offline_since, talk_hash FROM endpoints WHERE id > ? - AND (lower(id) LIKE ? OR lower(name) LIKE ?) + AND (lower(id) LIKE ? ESCAPE '\' OR lower(name) LIKE ? ESCAPE '\') ORDER BY id ASC LIMIT ?`, cursor, prefix, like, limit+1) } diff --git a/internal/app/presence/like.go b/internal/app/presence/like.go new file mode 100644 index 0000000..72bec63 --- /dev/null +++ b/internal/app/presence/like.go @@ -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() +} diff --git a/internal/app/presence/presence_test.go b/internal/app/presence/presence_test.go index a46e5e0..8cafcfc 100644 --- a/internal/app/presence/presence_test.go +++ b/internal/app/presence/presence_test.go @@ -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) { t.Parallel() app, db, down := openPresence(t) diff --git a/internal/auth/locks.go b/internal/auth/locks.go index 2c97074..b2a4dd9 100644 --- a/internal/auth/locks.go +++ b/internal/auth/locks.go @@ -21,23 +21,34 @@ func policyFor(kind LockKind) lockPolicy { } } +const ( + defaultMaxLockEntries = 65536 + lockSweepEveryOps = 1024 + lockSweepInterval = time.Minute +) + type lockEntry struct { fails []time.Time lockedUntil time.Time + lastFail time.Time } // MemoryLocks 是内存锁定计数器(重启清零)。 type MemoryLocks struct { - mu sync.Mutex - entries map[string]*lockEntry - now func() time.Time + mu sync.Mutex + entries map[string]*lockEntry + now func() time.Time + ops int + lastSweep time.Time + maxEntries int } // NewLoginLocks 创建默认锁定计数器。 func NewLoginLocks() *MemoryLocks { return &MemoryLocks{ - entries: make(map[string]*lockEntry), - now: time.Now, + entries: make(map[string]*lockEntry), + now: time.Now, + maxEntries: defaultMaxLockEntries, } } @@ -56,23 +67,121 @@ func lockMapKey(key LockKey) string { 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 与剩余时间。 func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) { l.mu.Lock() defer l.mu.Unlock() now := l.now() - e := l.entries[lockMapKey(key)] + k := lockMapKey(key) + e := l.entries[k] if e == nil { + l.maybeSweepLocked(now) return false, 0 } if e.lockedUntil.After(now) { + l.maybeSweepLocked(now) return true, e.lockedUntil.Sub(now) } - // 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。 - if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) { + if e.stale(now, policyFor(key.Kind).window) { + delete(l.entries, k) + l.maybeSweepLocked(now) + return false, 0 + } + // 到期自动解除:清空锁定与窗口内失败(保留仍在窗口内的失败计数)。 + if !e.lockedUntil.IsZero() { e.lockedUntil = time.Time{} e.fails = nil } + l.maybeSweepLocked(now) return false, 0 } @@ -88,6 +197,7 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) { l.entries[k] = e } if e.lockedUntil.After(now) { + l.maybeSweepLocked(now) return true, e.lockedUntil.Sub(now) } if !e.lockedUntil.IsZero() { @@ -103,11 +213,14 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) { } } e.fails = append(kept, now) + e.lastFail = now if len(e.fails) >= pol.maxFails { e.lockedUntil = now.Add(pol.lockFor) e.fails = nil + l.maybeSweepLocked(now) return true, pol.lockFor } + l.maybeSweepLocked(now) return false, 0 } diff --git a/internal/auth/locks_test.go b/internal/auth/locks_test.go new file mode 100644 index 0000000..7be4ce9 --- /dev/null +++ b/internal/auth/locks_test.go @@ -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) + } +} diff --git a/web/e2e/admin-main.spec.ts b/web/e2e/admin-main.spec.ts index 6ce66f9..d862b77 100644 --- a/web/e2e/admin-main.spec.ts +++ b/web/e2e/admin-main.spec.ts @@ -106,11 +106,12 @@ test.describe("后台主路径 W4", () => { await nav(page, "注册设置"); await expect(page).toHaveURL(/\/registration/); - await page.getByTestId("reg-enabled").click(); - await expect(page.getByText("已开启自助注册")).toBeVisible(); + await expect(page.getByTestId("reg-code-missing")).toBeVisible(); await page.getByTestId("reg-code").locator("input").fill("w4-reg-code-01"); await page.getByTestId("reg-save-code").click(); await expect(page.getByText("安全码已更新")).toBeVisible(); + await page.getByTestId("reg-enabled").click(); + await expect(page.getByText("已开启自助注册")).toBeVisible(); await nav(page, "群"); await expect(page).toHaveURL(/\/groups/); diff --git a/web/src/views/RegistrationView.spec.ts b/web/src/views/RegistrationView.spec.ts new file mode 100644 index 0000000..6e436b0 --- /dev/null +++ b/web/src/views/RegistrationView.spec.ts @@ -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(); + }); +}); diff --git a/web/src/views/RegistrationView.vue b/web/src/views/RegistrationView.vue index 7a4a44c..ce4067c 100644 --- a/web/src/views/RegistrationView.vue +++ b/web/src/views/RegistrationView.vue @@ -1,6 +1,6 @@