From 5176e04995dda09a0be571b02648c1cd1999e0dc Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 14:55:36 +0800 Subject: [PATCH 01/11] =?UTF-8?q?fix:=20=E9=99=90=E5=88=B6=E7=AE=A1?= =?UTF-8?q?=E7=90=86=E6=8E=A5=E5=8F=A3=E8=AF=B7=E6=B1=82=E4=BD=93=E5=A4=A7?= =?UTF-8?q?=E5=B0=8F=E4=B8=8E=E8=AF=BB=E5=8F=96=E6=97=B6=E9=97=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 ++++++ internal/admin/endpoints_csv.go | 17 +++++++++-- internal/admin/h01_test.go | 53 +++++++++++++++++++++++++++++++++ internal/admin/handler.go | 36 +++++++++++++++++++++- internal/admin/login.go | 4 +-- internal/httpx/json.go | 20 +++++++++++-- internal/httpx/json_test.go | 39 ++++++++++++++++++++++++ 7 files changed, 170 insertions(+), 8 deletions(-) create mode 100644 internal/admin/h01_test.go create mode 100644 internal/httpx/json_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index ac2fa7e..169e94a 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -694,6 +694,15 @@ - 备选:无。 - 影响:无。 +### 复审修复 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。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go index 11b3fa5..eff12ec 100644 --- a/internal/admin/endpoints_csv.go +++ b/internal/admin/endpoints_csv.go @@ -37,6 +37,11 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { raw, err := readImportCSV(r) if err != nil { + if httpx.IsBodyTooLarge(err) { + h.audit(actorString(p), "endpoint_import", "", "payload_too_large", ip) + httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大") + return + } h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error()) return @@ -91,9 +96,12 @@ func readImportCSV(r *http.Request) ([]byte, error) { } name := part.FormName() if name == "file" || name == "" { - b, readErr := io.ReadAll(io.LimitReader(part, 8<<20)) + b, readErr := io.ReadAll(part) _ = part.Close() if readErr != nil { + if httpx.IsBodyTooLarge(readErr) { + return nil, readErr + } return nil, errBadRequest("读取文件失败") } return b, nil @@ -102,9 +110,12 @@ func readImportCSV(r *http.Request) ([]byte, error) { } return nil, errBadRequest("缺少 file 字段") default: - // text/csv 或未标明时按原始体 - b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20)) + // text/csv 或未标明时按原始体;大小由 ServeHTTP 的 MaxBytesReader 限制。 + b, readErr := io.ReadAll(r.Body) if readErr != nil { + if httpx.IsBodyTooLarge(readErr) { + return nil, readErr + } return nil, errBadRequest("读取 CSV 失败") } return b, nil diff --git a/internal/admin/h01_test.go b/internal/admin/h01_test.go new file mode 100644 index 0000000..f1e11d4 --- /dev/null +++ b/internal/admin/h01_test.go @@ -0,0 +1,53 @@ +package admin_test + +import ( + "bytes" + "io" + "net/http" + "strings" + "testing" +) + +func TestLoginRejectsOversizeBody(t *testing.T) { + _, srv, client, _ := setup(t) + body := `{"username":"admin","password":"` + strings.Repeat("a", 1<<20) + `"}` + req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/login", strings.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + res, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + env := decodeEnv(t, res) + if res.StatusCode != http.StatusRequestEntityTooLarge { + t.Fatalf("want 413 got %d env=%+v", res.StatusCode, env) + } + if env.Error == nil || env.Error.Code != "payload_too_large" { + t.Fatalf("want payload_too_large got %+v", env.Error) + } +} + +func TestImportRejectsOversizeCSV(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + payload := bytes.Repeat([]byte("x"), 8<<20+1) + req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints/import", bytes.NewReader(payload)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "text/csv") + req.Header.Set("X-Nixmsg-Request", "1") + res, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + raw, _ := io.ReadAll(res.Body) + _ = res.Body.Close() + if res.StatusCode != http.StatusRequestEntityTooLarge { + t.Fatalf("want 413 got %d body=%s", res.StatusCode, raw) + } + if !strings.Contains(string(raw), "payload_too_large") { + t.Fatalf("want payload_too_large in body: %s", raw) + } +} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index e44b0a8..6787226 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -1,6 +1,7 @@ package admin import ( + "errors" "log/slog" "net" "net/http" @@ -11,6 +12,7 @@ import ( "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/config" + "git.asio.asia/nixevol/NixMsg/internal/httpx" "git.asio.asia/nixevol/NixMsg/internal/store" ) @@ -23,6 +25,13 @@ const ( defaultSessionTTL = 12 * time.Hour minPasswordLen = 12 lastUsedMinGap = time.Minute + + maxLoginBodyBytes = 8 << 10 + maxJSONBodyBytes = 1 << 20 + maxImportBodyBytes = 8 << 20 + defaultReadFor = 15 * time.Second + loginReadFor = 10 * time.Second + importReadFor = 2 * time.Minute ) // Deps 是管理 Handler 的依赖。 @@ -129,11 +138,36 @@ func New(d Deps) *Handler { return h } -// ServeHTTP 实现 http.Handler。 +// ServeHTTP 实现 http.Handler。按路由限制请求体大小与读截止时间。 func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + limit, readFor := requestBodyBudget(r) + if err := http.NewResponseController(w).SetReadDeadline(time.Now().Add(readFor)); err != nil && !errors.Is(err, http.ErrNotSupported) { + // 测试用 ResponseRecorder 或不支持截止时间的封装:忽略。 + } + if r.Body != nil { + r.Body = http.MaxBytesReader(w, r.Body, limit) + } h.mux.ServeHTTP(w, r) } +func requestBodyBudget(r *http.Request) (int64, time.Duration) { + if r.Method == http.MethodPost && r.URL.Path == "/api/admin/endpoints/import" { + return maxImportBodyBytes, importReadFor + } + if r.Method == http.MethodPost && r.URL.Path == "/api/admin/login" { + return maxLoginBodyBytes, loginReadFor + } + return maxJSONBodyBytes, defaultReadFor +} + +func writeDecodeError(w http.ResponseWriter, err error) { + if httpx.IsBodyTooLarge(err) { + httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大") + return + } + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") +} + func (h *Handler) routes() { // 公开 h.mux.HandleFunc("POST /api/admin/login", h.handleLogin) diff --git a/internal/admin/login.go b/internal/admin/login.go index 306184e..63897a0 100644 --- a/internal/admin/login.go +++ b/internal/admin/login.go @@ -35,7 +35,7 @@ func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) { Password string `json:"password"` } if err := httpx.DecodeJSON(r, &req); err != nil { - httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + writeDecodeError(w, err) return } if req.Username != adminUsername { @@ -117,7 +117,7 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { } if err := httpx.DecodeJSON(r, &req); err != nil { h.audit(actorString(p), "password_change", "", "bad_request", ip) - httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + writeDecodeError(w, err) return } if len(req.NewPassword) < minPasswordLen { diff --git a/internal/httpx/json.go b/internal/httpx/json.go index ecfac96..93175df 100644 --- a/internal/httpx/json.go +++ b/internal/httpx/json.go @@ -8,6 +8,17 @@ import ( "net/http" ) +const maxJSONBodyBytes = 1 << 20 + +// IsBodyTooLarge 判断是否因请求体超过 MaxBytesReader 上限而失败。 +func IsBodyTooLarge(err error) bool { + if err == nil { + return false + } + var maxErr *http.MaxBytesError + return errors.As(err, &maxErr) +} + // ErrorBody 是失败响应里的 error 对象。 type ErrorBody struct { Code string `json:"code"` @@ -46,10 +57,15 @@ func WriteError(w http.ResponseWriter, status int, code, message string) { }) } -// DecodeJSON 解码请求 JSON 体;空体对 dst 保持零值。 +// DecodeJSON 解码请求 JSON 体;空体对 dst 保持零值。内部把请求体限制在 1 MiB。 func DecodeJSON(r *http.Request, dst any) error { defer func() { _ = r.Body.Close() }() - dec := json.NewDecoder(r.Body) + body := r.Body + if body == nil { + return nil + } + body = http.MaxBytesReader(nil, body, maxJSONBodyBytes) + dec := json.NewDecoder(body) dec.DisallowUnknownFields() if err := dec.Decode(dst); err != nil { if errors.Is(err, io.EOF) { diff --git a/internal/httpx/json_test.go b/internal/httpx/json_test.go new file mode 100644 index 0000000..9307cc4 --- /dev/null +++ b/internal/httpx/json_test.go @@ -0,0 +1,39 @@ +package httpx + +import ( + "bytes" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestDecodeJSONRejectsOversizeBody(t *testing.T) { + t.Parallel() + payload := `{"password":"` + strings.Repeat("a", 2<<20) + `"}` + req := httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(payload)) + var dst struct { + Password string `json:"password"` + } + err := DecodeJSON(req, &dst) + if err == nil { + t.Fatal("want error for 2 MiB JSON body") + } + if !IsBodyTooLarge(err) { + t.Fatalf("want IsBodyTooLarge, got %v", err) + } +} + +func TestDecodeJSONAcceptsSmallBody(t *testing.T) { + t.Parallel() + req := httptest.NewRequest(http.MethodPost, "/x", bytes.NewReader([]byte(`{"password":"ok"}`))) + var dst struct { + Password string `json:"password"` + } + if err := DecodeJSON(req, &dst); err != nil { + t.Fatal(err) + } + if dst.Password != "ok" { + t.Fatalf("password=%q", dst.Password) + } +} From 512b8c2ba166e67733a0becb748815c378db9975 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 14:56:18 +0800 Subject: [PATCH 02/11] =?UTF-8?q?fix:=20=E7=AB=AF=E5=90=AF=E5=81=9C?= =?UTF-8?q?=E6=94=B9=E4=B8=BA=E4=BB=85=E7=94=B1=20identity=20=E6=89=A7?= =?UTF-8?q?=E8=A1=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 +++ internal/admin/endpoints.go | 7 +- internal/admin/h03_test.go | 146 ++++++++++++++++++++++++++++++++++++ 3 files changed, 161 insertions(+), 1 deletion(-) create mode 100644 internal/admin/h03_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 169e94a..13d674b 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -703,6 +703,15 @@ - 备选方案:登录 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 失败时端仍保持原启用状态。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index d6ef5bc..93c433e 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -350,8 +350,13 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) { hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil var wasEnabled bool var err error + // 注入 Identity 时启停只走 identity,避免先写 enabled 再级联失败造成半生效。 + patchEnabled := req.Enabled + if h.identity != nil { + patchEnabled = nil + } if hasMeta || (req.Enabled != nil && h.identity == nil) { - wasEnabled, err = h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, req.Enabled) + wasEnabled, err = h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, patchEnabled) if err != nil { if errors.Is(err, sql.ErrNoRows) { h.audit(actorString(p), "endpoint_patch", id, "not_found", ip) diff --git a/internal/admin/h03_test.go b/internal/admin/h03_test.go new file mode 100644 index 0000000..d4ebd92 --- /dev/null +++ b/internal/admin/h03_test.go @@ -0,0 +1,146 @@ +package admin_test + +import ( + "context" + "errors" + "net/http" + "sync" + "testing" + + "git.asio.asia/nixevol/NixMsg/internal/admin" + "git.asio.asia/nixevol/NixMsg/internal/app/identity" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" + "net/http/cookiejar" + "net/http/httptest" + "path/filepath" +) + +type failDisableIdentity struct { + identity.Stub + err error +} + +func (f *failDisableIdentity) Disable(context.Context, string) error { return f.err } +func (f *failDisableIdentity) Enable(context.Context, string) error { return nil } + +type trackIdentity struct { + identity.Stub + mu sync.Mutex + disableN int + enableN int +} + +func (t *trackIdentity) Disable(context.Context, string) error { + t.mu.Lock() + defer t.mu.Unlock() + t.disableN++ + return nil +} + +func (t *trackIdentity) Enable(context.Context, string) error { + t.mu.Lock() + defer t.mu.Unlock() + t.enableN++ + return nil +} + +func setupEndpointsWithIdentity(t *testing.T, ident identity.Service) (*store.DB, *httptest.Server, *http.Client) { + t.Helper() + dir := t.TempDir() + db, err := store.Open(filepath.Join(dir, "data"), "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + + hash := auth.NewStubHashPool() + if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil { + t.Fatal(seedErr) + } + h := admin.New(admin.Deps{ + DB: db, + Hash: hash, + Tokens: admin.NewRandomAPITokens(), + Locks: admin.NewMemoryLoginLocks(), + Identity: ident, + }) + srv := httptest.NewServer(h) + t.Cleanup(srv.Close) + + jar, err := cookiejar.New(nil) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Jar: jar} + res := postJSON(t, client, srv.URL+"/api/admin/login", + `{"username":"admin","password":"`+testPassword+`"}`, nil) + env := decodeEnv(t, res) + if res.StatusCode != http.StatusOK || !env.OK { + t.Fatalf("login failed: %d %+v", res.StatusCode, env) + } + return db, srv, client +} + +func TestPatchEnabledFalseKeepsEnabledWhenIdentityFails(t *testing.T) { + ident := &failDisableIdentity{err: errors.New("disable failed")} + db, srv, client := setupEndpointsWithIdentity(t, ident) + base := srv.URL + + res := postJSON(t, client, base+"/api/admin/endpoints", + `{"id":"keep-on","name":"仍启用","login_password":"password1"}`, + csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + + res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/keep-on", + `{"enabled":false}`, csrfHeaders()) + env = decodeEnv(t, res) + if res.StatusCode != http.StatusInternalServerError { + t.Fatalf("want 500 got %d %+v", res.StatusCode, env) + } + var enabled int + if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='keep-on'`).Scan(&enabled); err != nil { + t.Fatal(err) + } + if enabled != 1 { + t.Fatalf("want still enabled, got %d", enabled) + } +} + +func TestPatchNameOnlyDoesNotTouchEnabled(t *testing.T) { + ident := &trackIdentity{} + db, srv, client := setupEndpointsWithIdentity(t, ident) + base := srv.URL + + res := postJSON(t, client, base+"/api/admin/endpoints", + `{"id":"name-only","name":"旧名","login_password":"password1"}`, + csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + + res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/name-only", + `{"name":"新名"}`, csrfHeaders()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("patch: %d %+v", res.StatusCode, env) + } + ident.mu.Lock() + d, e := ident.disableN, ident.enableN + ident.mu.Unlock() + if d != 0 || e != 0 { + t.Fatalf("identity enable/disable should not run, disable=%d enable=%d", d, e) + } + var enabled int + var name string + if err := db.Read.QueryRow(`SELECT enabled, name FROM endpoints WHERE id='name-only'`).Scan(&enabled, &name); err != nil { + t.Fatal(err) + } + if enabled != 1 || name != "新名" { + t.Fatalf("enabled=%d name=%q", enabled, name) + } +} From 30330cea2c6cf34750e415b50f6becebaccff918 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 14:59:35 +0800 Subject: [PATCH 03/11] =?UTF-8?q?fix:=20=E6=89=B9=E9=87=8F=E5=AF=BC?= =?UTF-8?q?=E5=85=A5=E5=B9=B6=E5=8F=91=E5=93=88=E5=B8=8C=E5=B9=B6=E6=A0=A1?= =?UTF-8?q?=E6=AD=A3=20CSV=20=E8=A1=8C=E5=8F=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 ++ internal/admin/endpoints.go | 11 ++ internal/admin/endpoints_csv.go | 264 +++++++++++++++++++++--------- internal/admin/h04_test.go | 276 ++++++++++++++++++++++++++++++++ 4 files changed, 483 insertions(+), 77 deletions(-) create mode 100644 internal/admin/h04_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 13d674b..cfd741c 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -712,6 +712,15 @@ - 备选方案:同一事务里写 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。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index 93c433e..ba21f22 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -649,3 +649,14 @@ func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { Data: map[string]any{"errors": errs}, }) } + +func writeCSVConflictError(w http.ResponseWriter, errs []csvLineError) { + httpx.WriteJSON(w, http.StatusConflict, httpx.Envelope{ + OK: false, + Error: &httpx.ErrorBody{ + Code: "id_taken", + Message: "编号已占用", + }, + Data: map[string]any{"errors": errs}, + }) +} diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go index eff12ec..ee56630 100644 --- a/internal/admin/endpoints_csv.go +++ b/internal/admin/endpoints_csv.go @@ -5,12 +5,15 @@ import ( "context" "database/sql" "encoding/csv" + "errors" "io" "mime" "mime/multipart" "net/http" + "runtime" "strconv" "strings" + "sync" "unicode/utf8" "git.asio.asia/nixevol/NixMsg/internal/auth" @@ -47,7 +50,16 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { return } - prepared, errs := h.validateImportCSV(r.Context(), raw) + prepared, errs, fatal := h.validateImportCSV(r.Context(), raw) + if fatal != nil { + h.audit(actorString(p), "endpoint_import", "", "error", ip) + if errors.Is(fatal, context.Canceled) || errors.Is(fatal, context.DeadlineExceeded) { + httpx.WriteError(w, http.StatusServiceUnavailable, "busy", "哈希繁忙,请稍后再试") + return + } + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } if len(errs) > 0 { h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) writeCSVValidationError(w, errs) @@ -65,6 +77,11 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { }) } if err := h.insertEndpointsBatch(r.Context(), rows); err != nil { + if isUniqueConstraint(err) { + h.audit(actorString(p), "endpoint_import", "", "id_taken", ip) + writeCSVConflictError(w, h.importUniqueLineErrors(r.Context(), prepared)) + return + } h.audit(actorString(p), "endpoint_import", "", "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return @@ -128,59 +145,86 @@ func (e badRequestError) Error() string { return string(e) } func errBadRequest(msg string) error { return badRequestError(msg) } -func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError) { +type importPending struct { + line int + id string + name string + remark string + loginPW string + talkPW string + delaySec int64 + needGenerateID bool + needGenerateLogin bool +} + +func csvErrorLine(err error, fallback int) int { + var pe *csv.ParseError + if errors.As(err, &pe) { + if pe.StartLine > 0 { + return pe.StartLine + } + if pe.Line > 0 { + return pe.Line + } + } + if fallback > 0 { + return fallback + } + return 1 +} + +func importHashConcurrency() int { + n := runtime.NumCPU() - 1 + if n < 1 { + return 1 + } + return n +} + +func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError, error) { raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF}) reader := csv.NewReader(bytes.NewReader(raw)) reader.FieldsPerRecord = -1 reader.TrimLeadingSpace = true + reader.ReuseRecord = false - records, err := reader.ReadAll() + header, err := reader.Read() if err != nil { - return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}} + return nil, []csvLineError{{Line: csvErrorLine(err, 1), Reason: "CSV 解析失败"}}, nil } - if len(records) < 1 { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + headerLine, _ := reader.FieldPos(0) + if headerLine < 1 { + headerLine = 1 } - - header := normalizeCSVHeader(records[0]) + norm := normalizeCSVHeader(header) expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"} - if len(header) < len(expected) { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + if len(norm) < len(expected) { + return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil } for i, want := range expected { - if header[i] != want { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + if norm[i] != want { + return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil } } - dataRows := records[1:] - if len(dataRows) == 0 { - return nil, []csvLineError{{Line: 2, Reason: "没有数据行"}} - } - if len(dataRows) > maxImportRows { - return nil, []csvLineError{{Line: 1, Reason: "最多 1000 行"}} - } - errs := make([]csvLineError, 0) - prepared := make([]importPrepared, 0, len(dataRows)) - seen := make(map[string]int) // id -> first line - checkIDs := make([]string, 0, len(dataRows)) + seen := make(map[string]int) + checkIDs := make([]string, 0) + pendings := make([]importPending, 0) - type pending struct { - line int - id string - name string - remark string - loginPW string - talkPW string - delaySec int64 - needGenerateID bool - needGenerateLogin bool - } - pendings := make([]pending, 0, len(dataRows)) - - for i, cols := range dataRows { - line := i + 2 // 表头为 1 + for { + cols, readErr := reader.Read() + if errors.Is(readErr, io.EOF) { + break + } + if readErr != nil { + errs = append(errs, csvLineError{Line: csvErrorLine(readErr, 1), Reason: "CSV 解析失败"}) + continue + } + line, _ := reader.FieldPos(0) + if line < 1 { + line = headerLine + 1 + } for len(cols) < 6 { cols = append(cols, "") } @@ -209,7 +253,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr continue } - p := pending{ + p := importPending{ line: line, id: id, name: name, @@ -232,12 +276,18 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr } if len(errs) > 0 { - return nil, errs + return nil, errs, nil + } + if len(pendings) == 0 { + return nil, []csvLineError{{Line: headerLine + 1, Reason: "没有数据行"}}, nil + } + if len(pendings) > maxImportRows { + return nil, []csvLineError{{Line: headerLine, Reason: "最多 1000 行"}}, nil } existing, err := h.existingEndpointIDs(ctx, checkIDs) if err != nil { - return nil, []csvLineError{{Line: 1, Reason: "校验失败"}} + return nil, nil, err } for _, p := range pendings { if p.needGenerateID { @@ -248,16 +298,19 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr } } if len(errs) > 0 { - return nil, errs + return nil, errs, nil } - for _, p := range pendings { - useID := p.id + hashed := make([]importPending, len(pendings)) + copy(hashed, pendings) + for i := range hashed { + p := &hashed[i] if p.needGenerateID { + useID := "" for attempt := 0; attempt < 16; attempt++ { genID, genErr := generateEndpointID() if genErr != nil { - return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} + return nil, nil, genErr } if _, clash := seen[genID]; clash { continue @@ -270,46 +323,103 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr break } if useID == "" { - return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} + return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}}, nil } + p.id = useID } - - loginPW := p.loginPW if p.needGenerateLogin { pw, genErr := generateLoginPassword() if genErr != nil { - return nil, []csvLineError{{Line: p.line, Reason: "生成密码失败"}} + return nil, nil, genErr } - loginPW = pw + p.loginPW = pw } - loginHash, hashErr := h.hash.Hash(ctx, auth.PasswordLogin, loginPW) - if hashErr != nil { - return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}} - } - var talkHash sql.NullString - if p.talkPW != "" { - th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW) - if thErr != nil { - return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}} - } - talkHash = sql.NullString{String: th, Valid: true} - } - prepared = append(prepared, importPrepared{ - Insert: endpointInsert{ - ID: useID, - Name: p.name, - Remark: p.remark, - Source: sourceAdmin, - LoginHash: loginHash, - TalkHash: talkHash, - DefaultDelayMs: p.delaySec * 1000, - }, - PlainLogin: loginPW, - Name: p.name, - Line: p.line, - }) } - return prepared, nil + + prepared, hashErr := h.hashImportPendings(ctx, hashed) + if hashErr != nil { + return nil, nil, hashErr + } + return prepared, nil, nil +} + +func (h *Handler) hashImportPendings(ctx context.Context, pendings []importPending) ([]importPrepared, error) { + out := make([]importPrepared, len(pendings)) + errCh := make([]error, len(pendings)) + workers := importHashConcurrency() + sem := make(chan struct{}, workers) + var wg sync.WaitGroup + for i := range pendings { + i := i + wg.Add(1) + go func() { + defer wg.Done() + select { + case sem <- struct{}{}: + defer func() { <-sem }() + case <-ctx.Done(): + errCh[i] = ctx.Err() + return + } + p := pendings[i] + loginHash, hashErr := h.hash.Hash(ctx, auth.PasswordLogin, p.loginPW) + if hashErr != nil { + errCh[i] = hashErr + return + } + var talkHash sql.NullString + if p.talkPW != "" { + th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW) + if thErr != nil { + errCh[i] = thErr + return + } + talkHash = sql.NullString{String: th, Valid: true} + } + out[i] = importPrepared{ + Insert: endpointInsert{ + ID: p.id, + Name: p.name, + Remark: p.remark, + Source: sourceAdmin, + LoginHash: loginHash, + TalkHash: talkHash, + DefaultDelayMs: p.delaySec * 1000, + }, + PlainLogin: p.loginPW, + Name: p.name, + Line: p.line, + } + }() + } + wg.Wait() + for _, e := range errCh { + if e != nil { + return nil, e + } + } + return out, nil +} + +func (h *Handler) importUniqueLineErrors(ctx context.Context, prepared []importPrepared) []csvLineError { + ids := make([]string, 0, len(prepared)) + for _, p := range prepared { + ids = append(ids, p.Insert.ID) + } + existing, err := h.existingEndpointIDs(ctx, ids) + if err != nil { + return []csvLineError{{Line: 1, Reason: "编号已占用"}} + } + errs := make([]csvLineError, 0) + for _, p := range prepared { + if _, ok := existing[p.Insert.ID]; ok { + errs = append(errs, csvLineError{Line: p.Line, Reason: "编号已占用"}) + } + } + if len(errs) == 0 { + return []csvLineError{{Line: 1, Reason: "编号已占用"}} + } + return errs } func normalizeCSVHeader(cols []string) []string { diff --git a/internal/admin/h04_test.go b/internal/admin/h04_test.go new file mode 100644 index 0000000..a72e853 --- /dev/null +++ b/internal/admin/h04_test.go @@ -0,0 +1,276 @@ +package admin_test + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "path/filepath" + "runtime" + "strings" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/admin" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +type concHashPool struct { + inner *auth.StubHashPool + mu sync.Mutex + active, max int +} + +func (p *concHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) { + p.mu.Lock() + p.active++ + if p.active > p.max { + p.max = p.active + } + p.mu.Unlock() + time.Sleep(40 * time.Millisecond) + defer func() { + p.mu.Lock() + p.active-- + p.mu.Unlock() + }() + return p.inner.Hash(ctx, kind, password) +} + +func (p *concHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) { + return p.inner.Verify(ctx, kind, password, phc) +} + +func (p *concHashPool) QueueLen() int { return 0 } + +func (p *concHashPool) maxActive() int { + p.mu.Lock() + defer p.mu.Unlock() + return p.max +} + +type gateHashPool struct { + inner *auth.StubHashPool + started chan struct{} + release chan struct{} + startOnce sync.Once +} + +func (p *gateHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) { + p.startOnce.Do(func() { close(p.started) }) + select { + case <-p.release: + case <-ctx.Done(): + return "", ctx.Err() + } + return p.inner.Hash(ctx, kind, password) +} + +func (p *gateHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) { + return p.inner.Verify(ctx, kind, password, phc) +} + +func (p *gateHashPool) QueueLen() int { return 0 } + +func setupEndpointsHash(t *testing.T, hash auth.HashPool) (*store.DB, *httptest.Server, *http.Client) { + t.Helper() + dir := t.TempDir() + db, err := store.Open(filepath.Join(dir, "data"), "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + // 管理员口令用假哈希写入,避免导入用的 Hash 闸门卡住 seed。 + if seedErr := admin.SeedAdminPassword(context.Background(), db, auth.NewStubHashPool(), testPassword); seedErr != nil { + t.Fatal(seedErr) + } + h := admin.New(admin.Deps{ + DB: db, + Hash: hash, + Tokens: admin.NewRandomAPITokens(), + Locks: admin.NewMemoryLoginLocks(), + }) + srv := httptest.NewServer(h) + t.Cleanup(srv.Close) + jar, err := cookiejar.New(nil) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Jar: jar} + res := postJSON(t, client, srv.URL+"/api/admin/login", + `{"username":"admin","password":"`+testPassword+`"}`, nil) + env := decodeEnv(t, res) + if res.StatusCode != http.StatusOK || !env.OK { + t.Fatalf("login failed: %d %+v", res.StatusCode, env) + } + return db, srv, client +} + +func postImportCSV(t *testing.T, client *http.Client, base, csvBody string) *http.Response { + t.Helper() + req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/import", strings.NewReader(csvBody)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "text/csv") + req.Header.Set("X-Nixmsg-Request", "1") + res, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + return res +} + +func decodeCSVErrors(t *testing.T, res *http.Response) (status int, code string, lines []int) { + t.Helper() + raw, _ := io.ReadAll(res.Body) + _ = res.Body.Close() + var env struct { + OK bool `json:"ok"` + Error *struct { + Code string `json:"code"` + } `json:"error"` + Data *struct { + Errors []struct { + Line int `json:"line"` + Reason string `json:"reason"` + } `json:"errors"` + } `json:"data"` + } + if err := json.Unmarshal(raw, &env); err != nil { + t.Fatalf("json: %v body=%s", err, raw) + } + code = "" + if env.Error != nil { + code = env.Error.Code + } + if env.Data != nil { + for _, e := range env.Data.Errors { + lines = append(lines, e.Line) + } + } + return res.StatusCode, code, lines +} + +func TestImportHashConcurrencyGreaterThanOne(t *testing.T) { + if runtime.NumCPU() < 2 { + t.Skip("need at least 2 CPUs to observe concurrent hashing") + } + pool := &concHashPool{inner: auth.NewStubHashPool()} + _, srv, client := setupEndpointsHash(t, pool) + var b strings.Builder + b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") + for i := 0; i < 8; i++ { + fmt.Fprintf(&b, "conc-%d,名,password1,,0,\n", i) + } + res := postImportCSV(t, client, srv.URL, b.String()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("import: %d %+v", res.StatusCode, env) + } + if pool.maxActive() < 2 { + t.Fatalf("want concurrent hash > 1, got %d", pool.maxActive()) + } +} + +func TestImportThousandRows(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + var b strings.Builder + b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") + for i := 0; i < 1000; i++ { + fmt.Fprintf(&b, "r%04d,名,password1,,0,\n", i) + } + res := postImportCSV(t, client, srv.URL, b.String()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("import 1000: %d %+v", res.StatusCode, env) + } +} + +func TestImportUniqueRaceReturns409WithLine(t *testing.T) { + stub := auth.NewStubHashPool() + gate := &gateHashPool{inner: stub, started: make(chan struct{}), release: make(chan struct{})} + db, srv, client := setupEndpointsHash(t, gate) + csvBody := "" + + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + + "race-1,甲,password1,,0,\n" + + done := make(chan *http.Response, 1) + go func() { + done <- postImportCSV(t, client, srv.URL, csvBody) + }() + select { + case <-gate.started: + case <-time.After(5 * time.Second): + close(gate.release) + t.Fatal("hash did not start") + } + err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(` +INSERT INTO endpoints(id, name, remark, source, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) +VALUES ('race-1', '占', '', 'admin', 'stub$x', NULL, 0, 0, 1, 1)`) + return e + }) + if err != nil { + close(gate.release) + t.Fatal(err) + } + close(gate.release) + res := <-done + status, code, lines := decodeCSVErrors(t, res) + if status != http.StatusConflict || code != "id_taken" { + t.Fatalf("want 409 id_taken got %d %s lines=%v", status, code, lines) + } + found := false + for _, ln := range lines { + if ln == 2 { + found = true + } + } + if !found { + t.Fatalf("want line 2 in conflict errors, got %v", lines) + } +} + +func TestImportErrorLineSkipsEmptyRows(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + csvBody := "" + + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + + "\n" + + "bad-1,甲,password1,,-1,\n" + res := postImportCSV(t, client, srv.URL, csvBody) + status, _, lines := decodeCSVErrors(t, res) + if status != http.StatusBadRequest { + t.Fatalf("want 400 got %d lines=%v", status, lines) + } + found := false + for _, ln := range lines { + if ln == 3 { + found = true + } + } + if !found { + t.Fatalf("want physical line 3, got %v", lines) + } +} + +func TestImportUnclosedQuoteLine(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + csvBody := "" + + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + + "\"not-closed\n" + res := postImportCSV(t, client, srv.URL, csvBody) + status, _, lines := decodeCSVErrors(t, res) + if status != http.StatusBadRequest { + t.Fatalf("want 400 got %d lines=%v", status, lines) + } + if len(lines) == 0 || lines[0] < 2 { + t.Fatalf("want parse error line >= 2, got %v", lines) + } +} From ad7c50d5040dde544467a830c1af3b2566c936e7 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:05:37 +0800 Subject: [PATCH 04/11] =?UTF-8?q?fix:=20=E6=A0=A1=E9=AA=8C=E7=AB=AF?= =?UTF-8?q?=E9=BB=98=E8=AE=A4=E5=BB=B6=E8=BF=9F=E4=B8=8D=E8=B6=85=E8=BF=87?= =?UTF-8?q?=E8=B0=83=E5=BA=A6=E4=B8=8A=E9=99=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 +++++ internal/admin/endpoints.go | 28 +++++++++++---- internal/admin/endpoints_csv.go | 2 +- internal/admin/h05_test.go | 61 +++++++++++++++++++++++++++++++++ 4 files changed, 93 insertions(+), 7 deletions(-) create mode 100644 internal/admin/h05_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index cfd741c..397efef 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -721,6 +721,15 @@ - 备选方案:占满哈希池;未采用,给 MQTT 登录留槽位。 - 影响:导入吞吐上升;冲突不再表现为无行号的 500。 +### 复审修复 H-05 + +- 日期:2026-09-30 +- 原条款:PRD F11 最远可定到 365 天后(可配置);审查 #48。 +- 实际做法:`validateEndpointFields` 增加 `MaxScheduleSeconds`,开通/修改/导入与端侧改默认延迟同一上限(0 ≤ delay ≤ max)。界面 `:max` 在 W-05。 +- 原因:后台原先只检查 `>= 0`,可设出永远发不出消息的端,极大值还会溢出。 +- 备选方案:导入超限静默改成上限;未采用,改为行错误让管理员改表。 +- 影响:超上限返回 400 或 CSV 行错误。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index ba21f22..5e466cd 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -233,7 +233,7 @@ func (h *Handler) handleEndpointCreate(w http.ResponseWriter, r *http.Request) { if req.DefaultDelaySeconds != nil { delaySec = *req.DefaultDelaySeconds } - if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec); errMsg != "" { + if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec, h.maxScheduleSeconds()); errMsg != "" { h.audit(actorString(p), "endpoint_create", req.ID, "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg) return @@ -341,10 +341,12 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) { httpx.WriteError(w, http.StatusBadRequest, "bad_request", "备注过长") return } - if req.DefaultDelaySeconds != nil && *req.DefaultDelaySeconds < 0 { - h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) - httpx.WriteError(w, http.StatusBadRequest, "bad_request", "默认延迟无效") - return + if req.DefaultDelaySeconds != nil { + if msg := validateDelaySeconds(*req.DefaultDelaySeconds, h.maxScheduleSeconds()); msg != "" { + h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", msg) + return + } } hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil @@ -614,7 +616,7 @@ func (h *Handler) handleEndpointUnlock(w http.ResponseWriter, r *http.Request) { httpx.WriteOK(w, map[string]any{}) } -func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec int64) string { +func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec, maxDelaySec int64) string { if id != "" && !protocol.ValidEndpointID(id) { return "编号不合法" } @@ -633,12 +635,26 @@ func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec i if !protocol.ValidTalkPassword(talkPW) { return "对话密码不合法" } + if msg := validateDelaySeconds(delaySec, maxDelaySec); msg != "" { + return msg + } + return "" +} + +func validateDelaySeconds(delaySec, maxDelaySec int64) string { if delaySec < 0 { return "默认延迟无效" } + if maxDelaySec > 0 && delaySec > maxDelaySec { + return "默认延迟无效" + } return "" } +func (h *Handler) maxScheduleSeconds() int64 { + return int64(h.cfg.Limits.MaxScheduleSeconds) +} + func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{ OK: false, diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go index ee56630..bf9844f 100644 --- a/internal/admin/endpoints_csv.go +++ b/internal/admin/endpoints_csv.go @@ -244,7 +244,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr } delaySec = n } - if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec); msg != "" { + if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec, h.maxScheduleSeconds()); msg != "" { errs = append(errs, csvLineError{Line: line, Reason: msg}) continue } diff --git a/internal/admin/h05_test.go b/internal/admin/h05_test.go new file mode 100644 index 0000000..1af4142 --- /dev/null +++ b/internal/admin/h05_test.go @@ -0,0 +1,61 @@ +package admin_test + +import ( + "encoding/json" + "io" + "net/http" + "testing" +) + +func TestCreateRejectsDefaultDelayOverMax(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + body := `{"id":"dly-c","name":"延迟","login_password":"password1","default_delay_seconds":31536001}` + res := postJSON(t, client, srv.URL+"/api/admin/endpoints", body, csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != http.StatusBadRequest { + t.Fatalf("want 400 got %d %+v", res.StatusCode, env) + } +} + +func TestPatchRejectsDefaultDelayOverMax(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + res := postJSON(t, client, srv.URL+"/api/admin/endpoints", + `{"id":"dly-p","name":"延迟","login_password":"password1"}`, csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + res = doReq(t, client, http.MethodPatch, srv.URL+"/api/admin/endpoints/dly-p", + `{"default_delay_seconds":31536001}`, csrfHeaders()) + env = decodeEnv(t, res) + if res.StatusCode != http.StatusBadRequest { + t.Fatalf("want 400 got %d %+v", res.StatusCode, env) + } +} + +func TestImportRejectsDefaultDelayOverMax(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + csvBody := "" + + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + + "dly-i,甲,password1,,31536001,\n" + res := postImportCSV(t, client, srv.URL, csvBody) + raw, _ := io.ReadAll(res.Body) + _ = res.Body.Close() + if res.StatusCode != http.StatusBadRequest { + t.Fatalf("want 400 got %d body=%s", res.StatusCode, raw) + } + var parsed struct { + Data *struct { + Errors []struct { + Line int `json:"line"` + Reason string `json:"reason"` + } `json:"errors"` + } `json:"data"` + } + if err := json.Unmarshal(raw, &parsed); err != nil { + t.Fatal(err) + } + if parsed.Data == nil || len(parsed.Data.Errors) == 0 { + t.Fatalf("want row error, body=%s", raw) + } +} From eff4bd53be24a2d12c54dc3f4b751ba43ab459a2 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:05:55 +0800 Subject: [PATCH 05/11] =?UTF-8?q?fix:=20=E6=94=B9=E5=AF=86=E8=AE=A1?= =?UTF-8?q?=E5=85=A5=E9=94=81=E5=AE=9A=E5=B9=B6=E6=B8=85=E7=90=86=E8=BF=87?= =?UTF-8?q?=E6=9C=9F=E7=AE=A1=E7=90=86=E5=91=98=E4=BC=9A=E8=AF=9D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 +++++ internal/admin/h06_test.go | 75 ++++++++++++++++++++++++++++++++++++++ internal/admin/login.go | 23 +++++++++++- internal/admin/store.go | 3 ++ 4 files changed, 108 insertions(+), 2 deletions(-) create mode 100644 internal/admin/h06_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 397efef..e17c2ed 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -730,6 +730,15 @@ - 备选方案:导入超限静默改成上限;未采用,改为行错误让管理员改表。 - 影响:超上限返回 400 或 CSV 行错误。 +### 复审修复 H-06 + +- 日期:2026-09-30 +- 原条款:PRD §8 管理员登录防暴力;admin-api 2.4 改密作废其它会话;审查 #49。 +- 实际做法:改管理员密码时旧密码校验计入 `LockAdminIP`,锁定返回 429,成功清零。登录写入会话的同一事务删除过期 `admin_sessions`。作废其它会话失败返回 500 并写审计。 +- 原因:持有 Cookie 可无限猜旧密码;过期会话从不清理;作废失败被忽略。 +- 备选方案:改密失败不返回 401(与契约不符);未采用。W-01 排除改密 401 自动登出。 +- 影响:连续 10 次旧密码错误后锁定;过期会话行在下次登录清除。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/h06_test.go b/internal/admin/h06_test.go new file mode 100644 index 0000000..2d16503 --- /dev/null +++ b/internal/admin/h06_test.go @@ -0,0 +1,75 @@ +package admin_test + +import ( + "context" + "database/sql" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "path/filepath" + "testing" + + "git.asio.asia/nixevol/NixMsg/internal/admin" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +func TestPasswordWrongOldLocksAfterTen(t *testing.T) { + _, srv, client, _ := setup(t) + login(t, client, srv.URL) + base := srv.URL + var last *http.Response + var env envelope + for i := 0; i < 10; i++ { + last = postJSON(t, client, base+"/api/admin/password", + `{"old_password":"not-the-password","new_password":"new-password-12"}`, + map[string]string{"X-Nixmsg-Request": "1"}) + env = decodeEnv(t, last) + } + if last.StatusCode != http.StatusTooManyRequests { + t.Fatalf("want 429 after 10 wrong old passwords, got %d %+v", last.StatusCode, env) + } + if env.Error == nil || env.Error.Code != "rate_limited" { + t.Fatalf("want rate_limited got %+v", env.Error) + } +} + +func TestLoginDeletesExpiredAdminSessions(t *testing.T) { + dir := t.TempDir() + db, err := store.Open(filepath.Join(dir, "data"), "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + hash := auth.NewStubHashPool() + if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil { + t.Fatal(seedErr) + } + h := admin.New(admin.Deps{ + DB: db, + Hash: hash, + Tokens: admin.NewRandomAPITokens(), + Locks: admin.NewMemoryLoginLocks(), + }) + srv := httptest.NewServer(h) + t.Cleanup(srv.Close) + if err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES ('expired-hash', 1, 1)`) + return e + }); err != nil { + t.Fatal(err) + } + jar, err := cookiejar.New(nil) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Jar: jar} + login(t, client, srv.URL) + var n int + if err := db.Read.QueryRow(`SELECT COUNT(*) FROM admin_sessions WHERE token_hash = 'expired-hash'`).Scan(&n); err != nil { + t.Fatal(err) + } + if n != 0 { + t.Fatalf("expired session still present, count=%d", n) + } +} diff --git a/internal/admin/login.go b/internal/admin/login.go index 63897a0..972a732 100644 --- a/internal/admin/login.go +++ b/internal/admin/login.go @@ -111,6 +111,13 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { p, _ := principalFrom(r.Context()) ip := httpx.ClientIP(r, h.trusted) + if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked { + w.Header().Set("Retry-After", formatRetryAfter(retry)) + h.audit(actorString(p), "password_change", "", "rate_limited", ip) + httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试") + return + } + var req struct { OldPassword string `json:"old_password"` NewPassword string `json:"new_password"` @@ -134,10 +141,18 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { } ok, err := h.hash.Verify(r.Context(), auth.PasswordAdmin, req.OldPassword, phc) if err != nil || !ok { + locked, retry := h.locks.Fail(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}) + if locked { + w.Header().Set("Retry-After", formatRetryAfter(retry)) + h.audit(actorString(p), "password_change", "", "rate_limited", ip) + httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试") + return + } h.audit(actorString(p), "password_change", "", "unauthorized", ip) httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "旧密码错误") return } + h.locks.Clear(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}) newPHC, err := h.hash.Hash(r.Context(), auth.PasswordAdmin, req.NewPassword) if err != nil { h.audit(actorString(p), "password_change", "", "error", ip) @@ -149,9 +164,13 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } - // 保留当前会话,作废其它会话 + // 保留当前会话,作废其它会话;失败则返回 500,避免其它会话继续有效。 if p.Session != "" { - _ = h.deleteOtherSessions(r.Context(), hashSessionHex(p.Session)) + if err := h.deleteOtherSessions(r.Context(), hashSessionHex(p.Session)); err != nil { + h.audit(actorString(p), "password_change", "", "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } } h.audit(actorString(p), "password_change", "", "ok", ip) httpx.WriteOK(w, map[string]any{}) diff --git a/internal/admin/store.go b/internal/admin/store.go index d4e52b9..b736999 100644 --- a/internal/admin/store.go +++ b/internal/admin/store.go @@ -42,6 +42,9 @@ func (h *Handler) setAdminPasswordHash(ctx context.Context, phc string) error { func (h *Handler) createSession(ctx context.Context, hashHex string, ttl time.Duration) error { now := time.Now() return h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + if _, err := tx.Exec(`DELETE FROM admin_sessions WHERE expires_at <= ?`, now.UnixMilli()); err != nil { + return err + } _, err := tx.Exec( `INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES(?, ?, ?)`, hashHex, now.UnixMilli(), now.Add(ttl).UnixMilli(), From 898b9d2934c9d772f2eee4ed50bb47d8fa9ab323 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:12:16 +0800 Subject: [PATCH 06/11] =?UTF-8?q?fix:=20=E7=BB=9F=E4=B8=80=E7=AE=A1?= =?UTF-8?q?=E7=90=86=E5=90=8E=E5=8F=B0=20401=20=E4=B8=8E=E5=8A=A0=E8=BD=BD?= =?UTF-8?q?=E5=A4=B1=E8=B4=A5=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 10 ++++ internal/admin/login.go | 3 +- web/src/api/admin-mock.ts | 4 +- web/src/api/admin.ts | 4 +- web/src/api/http.spec.ts | 78 +++++++++++++++++++++++++++ web/src/api/http.ts | 86 +++++++++++++++++++++++++----- web/src/components/LoadFailed.vue | 19 +++++++ web/src/main.ts | 16 +++++- web/src/stores/auth.ts | 6 ++- web/src/views/EndpointsView.vue | 46 ++++++++++++---- web/src/views/GroupsView.vue | 8 +++ web/src/views/MessagesView.vue | 8 +++ web/src/views/OverviewView.spec.ts | 74 +++++++++++++++++++++++++ web/src/views/OverviewView.vue | 15 +++++- web/src/views/SettingsView.spec.ts | 60 +++++++++++++++++++++ web/src/views/SettingsView.vue | 75 +++++++++++++++++++++----- web/src/views/TokensView.vue | 8 +++ 17 files changed, 473 insertions(+), 47 deletions(-) create mode 100644 web/src/api/http.spec.ts create mode 100644 web/src/components/LoadFailed.vue create mode 100644 web/src/views/OverviewView.spec.ts create mode 100644 web/src/views/SettingsView.spec.ts diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index e17c2ed..4db93bc 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -797,6 +797,16 @@ - 备选:仅文档要求手跑 e2e;未采用。 - 影响:`task itest` 变长;需本机 Go/Node/Chromium。 +### 复审修复 W-01 + +- 日期:2026-09-30 +- 原条款:TASKS W1 请求封装;审查 #51;PRD F01 名称可选。 +- 实际做法:`requestAdmin` 对 401(排除 login/password)清空会话并跳转登录页,提示一次「登录已过期」;`run()` 不再对 `ApiError` 重复 toast;网络错误「无法连接服务器」、429「请稍后再试」。概览/系统/端/群/投递/令牌加载失败显示 `LoadFailed` 与重试。开通名称改为可选;改密用 n-form rules 按字符数校验,后端 `login.go` 同步 `utf8.RuneCountInString`;清除对话密码加确认。 +- 未改:`RegistrationView.vue` / `registration.go`(U-01 归属;本波指令禁止改注册页)。生成新安全码确认留给身份线。 +- 原因:会话过期后页面不跳转且错误弹两次;加载失败空白;帮助文字误称 nst_。 +- 备选方案:改密 401 也自动登出;未采用,与契约「旧密码错误 401」冲突。 +- 影响:改密失败不会踢当前会话;注册页错误提示仍走原逻辑。 + ### S1.1 传输层可注入假实现(Go / JS) - 相关文档:DEVELOPMENT 第 9 节单元测试要求「用假的 MQTT/HTTP,不要起真实服务器」。 diff --git a/internal/admin/login.go b/internal/admin/login.go index 972a732..6f8140f 100644 --- a/internal/admin/login.go +++ b/internal/admin/login.go @@ -4,6 +4,7 @@ import ( "database/sql" "errors" "net/http" + "unicode/utf8" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/httpx" @@ -127,7 +128,7 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { writeDecodeError(w, err) return } - if len(req.NewPassword) < minPasswordLen { + if utf8.RuneCountInString(req.NewPassword) < minPasswordLen { h.audit(actorString(p), "password_change", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", "新密码至少 12 位") return diff --git a/web/src/api/admin-mock.ts b/web/src/api/admin-mock.ts index 53a1f60..378006d 100644 --- a/web/src/api/admin-mock.ts +++ b/web/src/api/admin-mock.ts @@ -22,9 +22,7 @@ async function run(fn: () => Promise, silent = false): Promise { try { return await fn(); } catch (e) { - if (!silent && e instanceof ApiError) { - message.error(e.message); - } else if (!silent && e instanceof Error) { + if (!silent && e instanceof Error && !(e instanceof ApiError)) { message.error(e.message); } throw e; diff --git a/web/src/api/admin.ts b/web/src/api/admin.ts index 7944f69..99b1e20 100644 --- a/web/src/api/admin.ts +++ b/web/src/api/admin.ts @@ -33,9 +33,7 @@ async function run(fn: () => Promise, silent = false): Promise { try { return await fn(); } catch (e) { - if (!silent && e instanceof ApiError) { - message.error(e.message); - } else if (!silent && e instanceof Error) { + if (!silent && e instanceof Error && !(e instanceof ApiError)) { message.error(e.message); } throw e; diff --git a/web/src/api/http.spec.ts b/web/src/api/http.spec.ts new file mode 100644 index 0000000..afe7b5a --- /dev/null +++ b/web/src/api/http.spec.ts @@ -0,0 +1,78 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { ApiError, requestAdmin, setUnauthorizedHandler } from "./http"; + +vi.mock("@/utils/notify", () => ({ + message: { + error: vi.fn(), + success: vi.fn(), + warning: vi.fn(), + }, +})); + +import { message } from "@/utils/notify"; + +describe("requestAdmin", () => { + beforeEach(() => { + vi.stubGlobal("fetch", vi.fn()); + setUnauthorizedHandler(null); + vi.mocked(message.error).mockClear(); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + setUnauthorizedHandler(null); + }); + + it("401 跳转登录页且只提示一次", async () => { + const fetchMock = vi.mocked(fetch); + fetchMock.mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "unauthorized", message: "未登录" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + }), + ); + const handler = vi.fn(); + setUnauthorizedHandler(handler); + + await expect(requestAdmin("/api/admin/overview")).rejects.toBeInstanceOf(ApiError); + await expect(requestAdmin("/api/admin/endpoints")).rejects.toBeInstanceOf(ApiError); + + expect(handler).toHaveBeenCalled(); + expect(vi.mocked(message.error).mock.calls.filter((c) => c[0] === "登录已过期")).toHaveLength(1); + }); + + it("改密 401 不自动登出", async () => { + const fetchMock = vi.mocked(fetch); + fetchMock.mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "unauthorized", message: "旧密码错误" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + }), + ); + const handler = vi.fn(); + setUnauthorizedHandler(handler); + await expect( + requestAdmin("/api/admin/password", { method: "POST", body: { old_password: "x", new_password: "yyyyyyyyyyyy" } }), + ).rejects.toMatchObject({ message: "旧密码错误" }); + expect(handler).not.toHaveBeenCalled(); + }); + + it("网络错误显示中文", async () => { + vi.mocked(fetch).mockRejectedValue(new TypeError("Failed to fetch")); + await expect(requestAdmin("/api/admin/me")).rejects.toMatchObject({ message: "无法连接服务器" }); + expect(message.error).toHaveBeenCalledWith("无法连接服务器"); + }); + + it("429 提示稍后再试", async () => { + vi.mocked(fetch).mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "rate_limited", message: "locked" } }), { + status: 429, + headers: { "Content-Type": "application/json" }, + }), + ); + await expect(requestAdmin("/api/admin/login", { method: "POST", body: {} })).rejects.toMatchObject({ + message: "请求过于频繁,请稍后再试", + }); + expect(message.error).toHaveBeenCalledWith("请求过于频繁,请稍后再试"); + }); +}); diff --git a/web/src/api/http.ts b/web/src/api/http.ts index b92c1a4..6db2679 100644 --- a/web/src/api/http.ts +++ b/web/src/api/http.ts @@ -27,9 +27,52 @@ export interface RequestOptions { contentType?: string | null; } +const noAutoLogout = new Set(["/api/admin/login", "/api/admin/password"]); + +type UnauthorizedHandler = (redirectPath: string) => void; +let unauthorizedHandler: UnauthorizedHandler | null = null; +let expiredNotified = false; + +export function setUnauthorizedHandler(handler: UnauthorizedHandler | null) { + unauthorizedHandler = handler; +} + +function requestPath(path: string): string { + const q = path.indexOf("?"); + return q >= 0 ? path.slice(0, q) : path; +} + +function currentRedirect(): string { + if (typeof window === "undefined") { + return "/"; + } + return `${window.location.pathname}${window.location.search}` || "/"; +} + +function notifyExpiredOnce() { + if (expiredNotified) { + return; + } + expiredNotified = true; + message.error("登录已过期"); + window.setTimeout(() => { + expiredNotified = false; + }, 2000); +} + +function handleUnauthorized(path: string, silent: boolean) { + if (noAutoLogout.has(requestPath(path))) { + return; + } + if (!silent) { + notifyExpiredOnce(); + } + unauthorizedHandler?.(currentRedirect()); +} + /** - * 管理接口请求封装:自动带 credentials 与 X-Nixmsg-Request: 1,错误直接 message 显示。 - * W4 接真实后端时页面不必改,只需让 admin 模块走本函数。 + * 管理接口请求封装:自动带 credentials 与 X-Nixmsg-Request: 1。 + * 错误只在这里提示一次;401(登录/改密除外)清空会话并跳转登录页。 */ export async function requestAdmin(path: string, opts: RequestOptions = {}): Promise { const method = opts.method ?? (opts.body != null || opts.rawBody != null ? "POST" : "GET"); @@ -46,19 +89,32 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): headers["Content-Type"] = opts.contentType; } - const res = await fetch(path, { - method, - credentials: "include", - headers, - body: body ?? undefined, - }); + let res: Response; + try { + res = await fetch(path, { + method, + credentials: "include", + headers, + body: body ?? undefined, + }); + } catch { + const err = new ApiError("network", "无法连接服务器", 0); + if (!opts.silent) { + message.error(err.message); + } + throw err; + } + + if (res.status === 401) { + handleUnauthorized(path, Boolean(opts.silent)); + } let envelope: ApiEnvelope; try { envelope = (await res.json()) as ApiEnvelope; } catch { const err = new ApiError("internal", `响应不是 JSON(HTTP ${res.status})`, res.status); - if (!opts.silent) { + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; @@ -66,7 +122,7 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): if (typeof envelope !== "object" || envelope === null) { const err = new ApiError("internal", "响应格式无效", res.status); - if (!opts.silent) { + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; @@ -77,8 +133,14 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): } const errBody: ApiErrorBody = envelope.error ?? { code: "internal", message: "未知错误" }; - const err = new ApiError(errBody.code, errBody.message, res.status, "data" in envelope ? envelope.data : undefined); - if (!opts.silent) { + let display = errBody.message; + if (res.status === 429) { + display = "请求过于频繁,请稍后再试"; + } else if (res.status === 401 && !noAutoLogout.has(requestPath(path))) { + display = "登录已过期"; + } + const err = new ApiError(errBody.code, display, res.status, "data" in envelope ? envelope.data : undefined); + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; diff --git a/web/src/components/LoadFailed.vue b/web/src/components/LoadFailed.vue new file mode 100644 index 0000000..fd8b4e6 --- /dev/null +++ b/web/src/components/LoadFailed.vue @@ -0,0 +1,19 @@ + + + diff --git a/web/src/main.ts b/web/src/main.ts index 8c01e94..93fea8a 100644 --- a/web/src/main.ts +++ b/web/src/main.ts @@ -2,8 +2,22 @@ import { createApp } from "vue"; import { createPinia } from "pinia"; import App from "./App.vue"; import { router } from "./router"; +import { setUnauthorizedHandler } from "./api/http"; +import { useAuthStore } from "./stores/auth"; const app = createApp(App); -app.use(createPinia()); +const pinia = createPinia(); +app.use(pinia); 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"); diff --git a/web/src/stores/auth.ts b/web/src/stores/auth.ts index 7d35d41..82c08f7 100644 --- a/web/src/stores/auth.ts +++ b/web/src/stores/auth.ts @@ -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 }; }); diff --git a/web/src/views/EndpointsView.vue b/web/src/views/EndpointsView.vue index da4394e..ca93cc0 100644 --- a/web/src/views/EndpointsView.vue +++ b/web/src/views/EndpointsView.vue @@ -2,6 +2,7 @@ import { computed, h, onMounted, reactive, ref } from "vue"; import type { DataTableColumns, DataTableRowKey } from "naive-ui"; import { + NAlert, NButton, NDataTable, NForm, @@ -21,6 +22,7 @@ import type { UploadCustomRequestOptions } from "naive-ui"; import PageHeader from "@/components/PageHeader.vue"; import HelpTip from "@/components/HelpTip.vue"; import SecretOnceAlert from "@/components/SecretOnceAlert.vue"; +import LoadFailed from "@/components/LoadFailed.vue"; import { batchEndpoints, createEndpoint, @@ -41,6 +43,7 @@ import { message } from "@/utils/notify"; const dialog = useDialog(); const loading = ref(false); +const loadError = ref(""); const rows = ref([]); const total = ref(0); const checkedKeys = ref([]); @@ -77,6 +80,7 @@ const createForm = reactive({ default_delay_seconds: 0, }); const createLoading = ref(false); +const createError = ref(""); const editOpen = ref(false); const editId = ref(""); @@ -94,6 +98,7 @@ const talkPassword = ref(""); async function load() { loading.value = true; + loadError.value = ""; try { const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : ""; const res = await listEndpoints({ @@ -105,6 +110,8 @@ async function load() { }); rows.value = res.items; total.value = res.total; + } catch (e) { + loadError.value = e instanceof Error ? e.message : "加载失败"; } finally { loading.value = false; } @@ -187,14 +194,12 @@ function openCreate() { talk_password: "", default_delay_seconds: 0, }); + createError.value = ""; createOpen.value = true; } async function submitCreate() { - if (!createForm.name.trim()) { - message.error("名称不能为空"); - return; - } + createError.value = ""; createLoading.value = true; try { const res = await createEndpoint({ @@ -212,6 +217,8 @@ async function submitCreate() { filename: `${res.id}-password.txt`, }; await load(); + } catch (e) { + createError.value = e instanceof ApiError ? e.message : e instanceof Error ? e.message : "开通失败"; } finally { createLoading.value = false; } @@ -252,10 +259,23 @@ function openTalk(r: Endpoint) { } async function submitTalk() { - await setTalkPassword(talkId.value, talkPassword.value); - talkOpen.value = false; - message.success(talkPassword.value ? "已设置对话密码" : "已清除对话密码"); - await load(); + const save = async () => { + await setTalkPassword(talkId.value, talkPassword.value); + talkOpen.value = false; + message.success(talkPassword.value ? "已设置对话密码" : "已清除对话密码"); + await load(); + }; + if (!talkPassword.value) { + dialog.warning({ + title: "清除对话密码", + content: "留空保存将清除该端的对话密码,确定继续?", + positiveText: "清除", + negativeText: "取消", + onPositiveClick: () => save(), + }); + return; + } + await save(); } async function onResetPwd(r: Endpoint) { @@ -360,6 +380,8 @@ function onPageSizeChange(s: number) {
+ +
- + + - + @@ -464,7 +488,7 @@ function onPageSizeChange(s: number) { - +
+ +
diff --git a/web/src/views/MessagesView.vue b/web/src/views/MessagesView.vue index aca82d4..91f0a22 100644 --- a/web/src/views/MessagesView.vue +++ b/web/src/views/MessagesView.vue @@ -14,10 +14,12 @@ import { } from "naive-ui"; import PageHeader from "@/components/PageHeader.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 { formatLocalMs } from "@/utils/time"; const loading = ref(false); +const loadError = ref(""); const rows = ref([]); const total = ref(0); const page = ref(1); @@ -41,6 +43,7 @@ const stateOptions = [ async function load() { loading.value = true; + loadError.value = ""; try { const cursor = page.value > 1 ? String((page.value - 1) * pageSize.value) : ""; const res = await listMessages({ @@ -53,6 +56,8 @@ async function load() { }); rows.value = res.items; total.value = res.total; + } catch (e) { + loadError.value = e instanceof Error ? e.message : "加载失败"; } finally { loading.value = false; } @@ -128,6 +133,8 @@ async function openDetail(seq: number) {
+ +
diff --git a/web/src/views/OverviewView.spec.ts b/web/src/views/OverviewView.spec.ts new file mode 100644 index 0000000..922b83b --- /dev/null +++ b/web/src/views/OverviewView.spec.ts @@ -0,0 +1,74 @@ +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()).not.toContain("加载失败"); + w.unmount(); + }); +}); diff --git a/web/src/views/OverviewView.vue b/web/src/views/OverviewView.vue index 4773d3c..482fc4c 100644 --- a/web/src/views/OverviewView.vue +++ b/web/src/views/OverviewView.vue @@ -3,18 +3,28 @@ import { onMounted, ref } from "vue"; import { NDescriptions, NDescriptionsItem, NSpin } from "naive-ui"; import PageHeader from "@/components/PageHeader.vue"; import HelpTip from "@/components/HelpTip.vue"; +import LoadFailed from "@/components/LoadFailed.vue"; import { fetchOverview, type Overview } from "@/api/admin"; import { formatLocalMs } from "@/utils/time"; const loading = ref(true); +const loadError = ref(""); const data = ref(null); -onMounted(async () => { +async function load() { + loading.value = true; + loadError.value = ""; try { data.value = await fetchOverview(); + } catch (e) { + loadError.value = e instanceof Error ? e.message : "加载失败"; } finally { loading.value = false; } +} + +onMounted(() => { + void load(); }); @@ -23,7 +33,8 @@ onMounted(async () => {
- + + {{ data.version }}