diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index ee3b21f..5b3292a 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1358,3 +1358,12 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:只关 listener 时 MQTT 客户端看不到规范的停机原因码。 - 备选方案:等 L-03 一并做(本线仍提供 broker API,避免监听线无法调用)。 - 影响:进程退出时端会收到 server shutting down;监听器 HTTP 优雅停机仍未做。 + +### 复审修复 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;须先设码或一次提交开启与码。 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..5e0743c 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 { diff --git a/internal/app/identity/register_test.go b/internal/app/identity/register_test.go index 2d992eb..2ad8efe 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") 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 @@