From 63b09beead3f9371e1882224b1e28f9cf46bcafe Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 07:18:11 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9E=E7=8E=B0=E7=AE=A1=E7=90=86?= =?UTF-8?q?=E5=90=8E=E5=8F=B0=E7=AB=AF=E7=AE=A1=E7=90=86=E6=8E=A5=E5=8F=A3?= =?UTF-8?q?=EF=BC=88A2=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 35 +- internal/admin/endpoints.go | 598 +++++++++++++++++++++++++++++++ internal/admin/endpoints_csv.go | 310 ++++++++++++++++ internal/admin/endpoints_db.go | 384 ++++++++++++++++++++ internal/admin/endpoints_test.go | 364 +++++++++++++++++++ internal/admin/handler.go | 30 +- 6 files changed, 1705 insertions(+), 16 deletions(-) create mode 100644 internal/admin/endpoints.go create mode 100644 internal/admin/endpoints_csv.go create mode 100644 internal/admin/endpoints_db.go create mode 100644 internal/admin/endpoints_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index cc70333..93a1f2a 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -524,13 +524,13 @@ - 备选方案:另加 INTEGER 列或把数字存成文本并在 JSON 里发数字。 - 影响:W 线类型应按 `string` 解析令牌 id。 -4. **尚未实现的管理路由(鉴权中间件已生效,业务返回 501)** +4. **A1 范围外路由当时返回 501(A2 已实现端管理)** - 原条款:DEVELOPMENT 第 8 节完整路由表。 - - 实际做法:已实现 `login`/`logout`/`me`/`password` 与 `/api/admin/tokens` 全套。以下路由经鉴权后返回 `501 not_implemented`(属 A2/A3): - `GET /overview`;`endpoints` 列表/开通/import/batch/详情/改/删/kick/reset-login-password/talk-password/unlock;`registration` GET/PUT;`groups` 全部(含 `GET /groups/{id}`);`messages` 列表与详情;`GET /settings`。 + - 实际做法(A1 当时):已实现 `login`/`logout`/`me`/`password` 与 `/api/admin/tokens` 全套;其余经鉴权后 `501`。 + - A2 起:端相关路由已改为真实现,见下节;仍为 501 的有 `overview`、`registration`、`groups`、`messages`、`settings`(属 A3)。 - 原因:A1 范围仅鉴权、令牌、操作日志。 - 备选方案:无。 - - 影响:W/集成测试在 A2/A3 前勿依赖这些业务响应。 + - 影响:W/集成测试在 A3 前勿依赖尚未实现的业务响应。 5. **操作日志用 `slog` 结构化字段** - 原条款:写结构化日志(操作者、动作、对象、结果、来源 IP)。 @@ -539,6 +539,33 @@ - 备选方案:独立 audit 表。 - 影响:日志采集需按 msg=`admin_audit` 过滤。 +### A2 2026-09-30 + +1. **停用/删除完整级联留给 I5** + - 原条款:PRD F01 / DEVELOPMENT 7.6:停用作废未送达消息与发送中消息;删除另含退群、群主转让/解散、清授权与回执/发出记录等。 + - 实际做法:A2 直接写库:停用立刻 `enabled=0` 并清空 `session_hash`;删除删 `endpoints` 行并清相关 `talk_grants`;二者均调用可注入的 `Deps.KickEndpoint` 踢连接。不作废投递/消息、不转让群主、不清理群成员与回执。 + - 原因:本分支尚无 I5;TASKS 允许先接现有存储并在偏差中写明。 + - 备选方案:阻塞等待 I5;或在 A 线内复制级联 SQL(易与 I5 重复冲突)。 + - 影响:合入 I5 后应由身份服务 `Disable`/`Delete` 接管级联;总控接线把 kick 钩子接到 broker。`KickEndpoint` 为 nil 时踢线为 no-op(单元测试可注入)。 + +2. **在线状态读库字段,不依赖 presence 服务** + - 原条款:列表含是否在线、最近上下线;可按在线筛选。 + - 实际做法:用 `endpoints.online_since` / `offline_since` 判定在线(`online_since` 非空且大于 `offline_since` 或后者为空);列表项带 `online` / `online_since_ms` / `offline_since_ms`。 + - 原因:presence/N3 未在本 Handler 注入;库字段是契约字段。 + - 备选方案:注入 `presence.Service.IsOnline`。 + - 影响:上下线列未由 N/I 写入前,列表会显示离线。 + +3. **`login_locked` 仅反映按编号的登录锁定** + - 原条款:列表含 `login_locked`。 + - 实际做法:查 `LockLoginEndpoint`;不枚举「编号+IP」锁定。`unlock` 仍调用 `ClearEndpoint` 清两种。 + - 原因:`LoginLocks` 接口无「是否任一 IP 锁定」查询。 + - 备选方案:扩展 locks 接口。 + - 影响:仅 IP 档锁定时列表可能仍显示未锁定,但 unlock 可解除。 + +4. **仍未挂载到 `cmd/nixmsg`** + - 同 A1:可挂载 Handler;接线与 `KickEndpoint` 注入留给总控/后续。 + - 影响:集成测试需自行挂载 Handler。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go new file mode 100644 index 0000000..371b3a0 --- /dev/null +++ b/internal/admin/endpoints.go @@ -0,0 +1,598 @@ +package admin + +import ( + "context" + "database/sql" + "errors" + "net/http" + "strconv" + "strings" + "unicode/utf8" + + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/httpx" + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +const ( + maxRemarkChars = 200 + defaultListLimit = 50 + maxListLimit = 200 + maxImportRows = 1000 + sourceAdmin = "admin" + loginPasswordOnceKey = "login_password" +) + +// EndpointKickFunc 只断开端当前连接;是否作废令牌由调用方在踢线前决定。 +// 返回 kicked=true 表示当时有连接被断开。nil 表示未接线(踢线为 no-op)。 +type EndpointKickFunc func(ctx context.Context, endpointID string) (kicked bool, err error) + +type endpointRow struct { + ID string + Name string + Remark string + Source string + Enabled bool + TalkHash sql.NullString + DefaultDelayMs int64 + CreatedAtMs int64 + OnlineSinceMs sql.NullInt64 + OfflineSinceMs sql.NullInt64 + SessionIssuedAt sql.NullInt64 + SessionUsedAt sql.NullInt64 +} + +func (e endpointRow) toAPI(loginLocked bool, detail bool) map[string]any { + online, onlineSince, offlineSince := presenceFields(e.OnlineSinceMs, e.OfflineSinceMs) + out := map[string]any{ + "id": e.ID, + "name": e.Name, + "remark": e.Remark, + "source": e.Source, + "enabled": e.Enabled, + "online": online, + "online_since_ms": onlineSince, + "offline_since_ms": offlineSince, + "talk_password_set": e.TalkHash.Valid && e.TalkHash.String != "", + "default_delay_ms": e.DefaultDelayMs, + "created_at_ms": e.CreatedAtMs, + "login_locked": loginLocked, + } + if detail { + var issued, used any + if e.SessionIssuedAt.Valid { + issued = e.SessionIssuedAt.Int64 + } + if e.SessionUsedAt.Valid { + used = e.SessionUsedAt.Int64 + } + out["session_issued_at_ms"] = issued + out["session_used_at_ms"] = used + } + return out +} + +func presenceFields(onlineSince, offlineSince sql.NullInt64) (online bool, onlineMs, offlineMs any) { + if onlineSince.Valid && (!offlineSince.Valid || onlineSince.Int64 > offlineSince.Int64) { + online = true + onlineMs = onlineSince.Int64 + if offlineSince.Valid { + offlineMs = offlineSince.Int64 + } + return online, onlineMs, offlineMs + } + if onlineSince.Valid { + onlineMs = onlineSince.Int64 + } + if offlineSince.Valid { + offlineMs = offlineSince.Int64 + } + return false, onlineMs, offlineMs +} + +func (h *Handler) endpointLoginLocked(id string) bool { + locked, _ := h.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: id}) + return locked +} + +func (h *Handler) kickEndpoint(ctx context.Context, id string) (bool, error) { + if h.kick == nil { + return false, nil + } + return h.kick(ctx, id) +} + +func (h *Handler) handleEndpointList(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + limit := defaultListLimit + if s := q.Get("limit"); s != "" { + n, err := strconv.Atoi(s) + if err != nil || n < 1 { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "limit 无效") + return + } + if n > maxListLimit { + n = maxListLimit + } + limit = n + } + offset := 0 + if c := q.Get("cursor"); c != "" { + n, err := strconv.Atoi(c) + if err != nil || n < 0 { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "cursor 无效") + return + } + offset = n + } + source := q.Get("source") + if source != "" && source != "admin" && source != "self" { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "source 无效") + return + } + var onlineFilter *bool + if s := q.Get("online"); s != "" { + switch s { + case "true": + v := true + onlineFilter = &v + case "false": + v := false + onlineFilter = &v + default: + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "online 无效") + return + } + } + var enabledFilter *bool + if s := q.Get("enabled"); s != "" { + switch s { + case "true": + v := true + enabledFilter = &v + case "false": + v := false + enabledFilter = &v + default: + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "enabled 无效") + return + } + } + query := strings.TrimSpace(q.Get("query")) + + total, err := h.countEndpoints(r.Context(), source, onlineFilter, enabledFilter, query) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + rows, err := h.listEndpoints(r.Context(), source, onlineFilter, enabledFilter, query, limit, offset) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + items := make([]map[string]any, 0, len(rows)) + for _, row := range rows { + items = append(items, row.toAPI(h.endpointLoginLocked(row.ID), false)) + } + next := "" + if offset+len(rows) < total { + next = strconv.Itoa(offset + len(rows)) + } + httpx.WriteOK(w, map[string]any{ + "items": items, + "next_cursor": next, + "total": total, + }) +} + +func (h *Handler) handleEndpointCreate(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + + var req struct { + ID string `json:"id"` + Name string `json:"name"` + Remark string `json:"remark"` + LoginPassword string `json:"login_password"` + TalkPassword string `json:"talk_password"` + DefaultDelaySeconds *int64 `json:"default_delay_seconds"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "endpoint_create", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + delaySec := int64(0) + if req.DefaultDelaySeconds != nil { + delaySec = *req.DefaultDelaySeconds + } + if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec); errMsg != "" { + h.audit(actorString(p), "endpoint_create", req.ID, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg) + return + } + + id := req.ID + loginPW := req.LoginPassword + pwGenerated := false + if loginPW == "" { + pw, err := generateLoginPassword() + if err != nil { + h.audit(actorString(p), "endpoint_create", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + loginPW = pw + pwGenerated = true + } + loginHash, err := h.hash.Hash(r.Context(), auth.PasswordLogin, loginPW) + if err != nil { + h.audit(actorString(p), "endpoint_create", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + var talkHash sql.NullString + if req.TalkPassword != "" { + th, hashErr := h.hash.Hash(r.Context(), auth.PasswordTalk, req.TalkPassword) + if hashErr != nil { + h.audit(actorString(p), "endpoint_create", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + talkHash = sql.NullString{String: th, Valid: true} + } + + createdID, err := h.insertEndpoint(r.Context(), endpointInsert{ + ID: id, + Name: req.Name, + Remark: req.Remark, + Source: sourceAdmin, + LoginHash: loginHash, + TalkHash: talkHash, + DefaultDelayMs: delaySec * 1000, + }) + if err != nil { + if isUniqueConstraint(err) { + h.audit(actorString(p), "endpoint_create", id, "id_taken", ip) + httpx.WriteError(w, http.StatusConflict, "id_taken", "编号已占用") + return + } + h.audit(actorString(p), "endpoint_create", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + h.audit(actorString(p), "endpoint_create", createdID, "ok", ip) + data := map[string]any{"id": createdID} + if pwGenerated { + data[loginPasswordOnceKey] = loginPW + } + httpx.WriteOK(w, data) +} + +func (h *Handler) handleEndpointGet(w http.ResponseWriter, r *http.Request) { + id := r.PathValue("id") + row, err := h.getEndpoint(r.Context(), id) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + httpx.WriteOK(w, row.toAPI(h.endpointLoginLocked(row.ID), true)) +} + +func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + Name *string `json:"name"` + Remark *string `json:"remark"` + DefaultDelaySeconds *int64 `json:"default_delay_seconds"` + Enabled *bool `json:"enabled"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if req.Name == nil && req.Remark == nil && req.DefaultDelaySeconds == nil && req.Enabled == nil { + h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "无更新字段") + return + } + if req.Name != nil && !protocol.ValidName(*req.Name) { + h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "名称不合法") + return + } + if req.Remark != nil && utf8.RuneCountInString(*req.Remark) > maxRemarkChars { + h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) + 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 + } + + wasEnabled, err := h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, req.Enabled) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + h.audit(actorString(p), "endpoint_patch", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + h.audit(actorString(p), "endpoint_patch", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if req.Enabled != nil && !*req.Enabled && wasEnabled { + _, _ = h.kickEndpoint(r.Context(), id) + } + row, err := h.getEndpoint(r.Context(), id) + if err != nil { + h.audit(actorString(p), "endpoint_patch", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + h.audit(actorString(p), "endpoint_patch", id, "ok", ip) + httpx.WriteOK(w, row.toAPI(h.endpointLoginLocked(row.ID), true)) +} + +func (h *Handler) handleEndpointDelete(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + ok, err := h.deleteEndpointBasic(r.Context(), id) + if err != nil { + h.audit(actorString(p), "endpoint_delete", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if !ok { + h.audit(actorString(p), "endpoint_delete", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + _, _ = h.kickEndpoint(r.Context(), id) + h.audit(actorString(p), "endpoint_delete", id, "ok", ip) + httpx.WriteOK(w, map[string]any{}) +} + +func (h *Handler) handleEndpointBatch(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + + var req struct { + IDs []string `json:"ids"` + Action string `json:"action"` + } + if err := httpx.DecodeJSON(r, &req); err != nil || len(req.IDs) == 0 { + h.audit(actorString(p), "endpoint_batch", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + switch req.Action { + case "disable", "enable", "delete": + default: + h.audit(actorString(p), "endpoint_batch", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "action 无效") + return + } + + okIDs := make([]string, 0, len(req.IDs)) + failed := make([]map[string]string, 0) + for _, id := range req.IDs { + var opErr error + var found bool + switch req.Action { + case "disable": + found, opErr = h.setEndpointEnabled(r.Context(), id, false) + if found && opErr == nil { + _, _ = h.kickEndpoint(r.Context(), id) + } + case "enable": + found, opErr = h.setEndpointEnabled(r.Context(), id, true) + case "delete": + found, opErr = h.deleteEndpointBasic(r.Context(), id) + if found && opErr == nil { + _, _ = h.kickEndpoint(r.Context(), id) + } + } + if opErr != nil { + failed = append(failed, map[string]string{"id": id, "code": "internal"}) + continue + } + if !found { + failed = append(failed, map[string]string{"id": id, "code": "not_found"}) + continue + } + okIDs = append(okIDs, id) + } + h.audit(actorString(p), "endpoint_batch_"+req.Action, strings.Join(okIDs, ","), "ok", ip) + httpx.WriteOK(w, map[string]any{"ok_ids": okIDs, "failed": failed}) +} + +func (h *Handler) handleEndpointKick(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + exists, err := h.endpointExists(r.Context(), id) + if err != nil { + h.audit(actorString(p), "endpoint_kick", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if !exists { + h.audit(actorString(p), "endpoint_kick", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + kicked, err := h.kickEndpoint(r.Context(), id) + if err != nil { + h.audit(actorString(p), "endpoint_kick", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + h.audit(actorString(p), "endpoint_kick", id, "ok", ip) + httpx.WriteOK(w, map[string]any{"kicked": kicked}) +} + +func (h *Handler) handleEndpointResetLoginPassword(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + LoginPassword string `json:"login_password"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "endpoint_reset_login_password", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if !protocol.ValidLoginPassword(req.LoginPassword) { + h.audit(actorString(p), "endpoint_reset_login_password", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "登录密码不合法") + return + } + + pw := req.LoginPassword + if pw == "" { + gen, err := generateLoginPassword() + if err != nil { + h.audit(actorString(p), "endpoint_reset_login_password", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + pw = gen + } + hash, err := h.hash.Hash(r.Context(), auth.PasswordLogin, pw) + if err != nil { + h.audit(actorString(p), "endpoint_reset_login_password", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + ok, err := h.resetLoginPassword(r.Context(), id, hash) + if err != nil { + h.audit(actorString(p), "endpoint_reset_login_password", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if !ok { + h.audit(actorString(p), "endpoint_reset_login_password", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + _, _ = h.kickEndpoint(r.Context(), id) + h.audit(actorString(p), "endpoint_reset_login_password", id, "ok", ip) + httpx.WriteOK(w, map[string]any{loginPasswordOnceKey: pw}) +} + +func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + TalkPassword string `json:"talk_password"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "endpoint_talk_password", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if !protocol.ValidTalkPassword(req.TalkPassword) { + h.audit(actorString(p), "endpoint_talk_password", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "对话密码不合法") + return + } + + var talkHash sql.NullString + if req.TalkPassword != "" { + th, err := h.hash.Hash(r.Context(), auth.PasswordTalk, req.TalkPassword) + if err != nil { + h.audit(actorString(p), "endpoint_talk_password", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + talkHash = sql.NullString{String: th, Valid: true} + } + ok, err := h.setTalkPassword(r.Context(), id, talkHash) + if err != nil { + h.audit(actorString(p), "endpoint_talk_password", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if !ok { + h.audit(actorString(p), "endpoint_talk_password", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + h.audit(actorString(p), "endpoint_talk_password", id, "ok", ip) + httpx.WriteOK(w, map[string]any{"talk_password_set": talkHash.Valid}) +} + +func (h *Handler) handleEndpointUnlock(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + exists, err := h.endpointExists(r.Context(), id) + if err != nil { + h.audit(actorString(p), "endpoint_unlock", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + if !exists { + h.audit(actorString(p), "endpoint_unlock", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") + return + } + h.locks.ClearEndpoint(id) + h.audit(actorString(p), "endpoint_unlock", id, "ok", ip) + httpx.WriteOK(w, map[string]any{}) +} + +func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec int64) string { + if id != "" && !protocol.ValidEndpointID(id) { + return "编号不合法" + } + if !protocol.ValidName(name) { + return "名称不合法" + } + if utf8.RuneCountInString(remark) > maxRemarkChars { + return "备注过长" + } + if protocol.LoginPasswordForbiddenPrefix(loginPW) { + return "登录密码不能以 nst_ 开头" + } + if !protocol.ValidLoginPassword(loginPW) { + return "登录密码不合法" + } + if !protocol.ValidTalkPassword(talkPW) { + return "对话密码不合法" + } + if delaySec < 0 { + return "默认延迟无效" + } + return "" +} + +func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { + httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{ + OK: false, + Error: &httpx.ErrorBody{ + Code: "bad_request", + Message: "CSV 校验失败", + }, + Data: map[string]any{"errors": errs}, + }) +} diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go new file mode 100644 index 0000000..11b3fa5 --- /dev/null +++ b/internal/admin/endpoints_csv.go @@ -0,0 +1,310 @@ +package admin + +import ( + "bytes" + "context" + "database/sql" + "encoding/csv" + "io" + "mime" + "mime/multipart" + "net/http" + "strconv" + "strings" + "unicode/utf8" + + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/httpx" + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +type csvLineError struct { + Line int `json:"line"` + Reason string `json:"reason"` +} + +type importPrepared struct { + Insert endpointInsert + // PlainLogin 始终填入响应(生成或原文,仅此一次)。 + PlainLogin string + Name string + Line int +} + +func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + + raw, err := readImportCSV(r) + if err != nil { + h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error()) + return + } + + prepared, errs := h.validateImportCSV(r.Context(), raw) + if len(errs) > 0 { + h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) + writeCSVValidationError(w, errs) + return + } + + rows := make([]endpointInsert, len(prepared)) + items := make([]map[string]any, 0, len(prepared)) + for i, pRow := range prepared { + rows[i] = pRow.Insert + items = append(items, map[string]any{ + "id": pRow.Insert.ID, + "login_password": pRow.PlainLogin, + "name": pRow.Name, + }) + } + if err := h.insertEndpointsBatch(r.Context(), rows); err != nil { + h.audit(actorString(p), "endpoint_import", "", "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + h.audit(actorString(p), "endpoint_import", strconv.Itoa(len(items)), "ok", ip) + httpx.WriteOK(w, map[string]any{"items": items}) +} + +func readImportCSV(r *http.Request) ([]byte, error) { + ct := r.Header.Get("Content-Type") + mediaType, params, err := mime.ParseMediaType(ct) + if err != nil { + mediaType = strings.TrimSpace(strings.Split(ct, ";")[0]) + } + switch { + case strings.HasPrefix(mediaType, "multipart/"): + boundary := params["boundary"] + if boundary == "" { + return nil, errBadRequest("multipart 缺少 boundary") + } + mr := multipart.NewReader(r.Body, boundary) + for { + part, pErr := mr.NextPart() + if pErr == io.EOF { + break + } + if pErr != nil { + return nil, errBadRequest("读取 multipart 失败") + } + name := part.FormName() + if name == "file" || name == "" { + b, readErr := io.ReadAll(io.LimitReader(part, 8<<20)) + _ = part.Close() + if readErr != nil { + return nil, errBadRequest("读取文件失败") + } + return b, nil + } + _ = part.Close() + } + return nil, errBadRequest("缺少 file 字段") + default: + // text/csv 或未标明时按原始体 + b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20)) + if readErr != nil { + return nil, errBadRequest("读取 CSV 失败") + } + return b, nil + } +} + +type badRequestError string + +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) { + raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF}) + reader := csv.NewReader(bytes.NewReader(raw)) + reader.FieldsPerRecord = -1 + reader.TrimLeadingSpace = true + + records, err := reader.ReadAll() + if err != nil { + return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}} + } + if len(records) < 1 { + return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + } + + header := normalizeCSVHeader(records[0]) + expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"} + if len(header) < len(expected) { + return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + } + for i, want := range expected { + if header[i] != want { + return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + } + } + + dataRows := records[1:] + if len(dataRows) == 0 { + return nil, []csvLineError{{Line: 2, Reason: "没有数据行"}} + } + if len(dataRows) > maxImportRows { + return nil, []csvLineError{{Line: 1, Reason: "最多 1000 行"}} + } + + errs := make([]csvLineError, 0) + prepared := make([]importPrepared, 0, len(dataRows)) + seen := make(map[string]int) // id -> first line + checkIDs := make([]string, 0, len(dataRows)) + + type pending struct { + line int + 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 len(cols) < 6 { + cols = append(cols, "") + } + id := strings.TrimSpace(cols[0]) + name := cols[1] + loginPW := cols[2] + talkPW := cols[3] + delayRaw := strings.TrimSpace(cols[4]) + remark := cols[5] + + delaySec := int64(0) + if delayRaw != "" { + n, parseErr := strconv.ParseInt(delayRaw, 10, 64) + if parseErr != nil || n < 0 { + errs = append(errs, csvLineError{Line: line, Reason: "默认延迟无效"}) + continue + } + delaySec = n + } + if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec); msg != "" { + errs = append(errs, csvLineError{Line: line, Reason: msg}) + continue + } + if utf8.RuneCountInString(name) > protocol.MaxNameChars { + errs = append(errs, csvLineError{Line: line, Reason: "名称不合法"}) + continue + } + + p := pending{ + line: line, + id: id, + name: name, + remark: remark, + loginPW: loginPW, + talkPW: talkPW, + delaySec: delaySec, + needGenerateID: id == "", + needGenerateLogin: loginPW == "", + } + if !p.needGenerateID { + if first, ok := seen[id]; ok { + errs = append(errs, csvLineError{Line: line, Reason: "编号与第 " + strconv.Itoa(first) + " 行重复"}) + continue + } + seen[id] = line + checkIDs = append(checkIDs, id) + } + pendings = append(pendings, p) + } + + if len(errs) > 0 { + return nil, errs + } + + existing, err := h.existingEndpointIDs(ctx, checkIDs) + if err != nil { + return nil, []csvLineError{{Line: 1, Reason: "校验失败"}} + } + for _, p := range pendings { + if p.needGenerateID { + continue + } + if _, ok := existing[p.id]; ok { + errs = append(errs, csvLineError{Line: p.line, Reason: "编号已占用"}) + } + } + if len(errs) > 0 { + return nil, errs + } + + for _, p := range pendings { + useID := p.id + if p.needGenerateID { + for attempt := 0; attempt < 16; attempt++ { + genID, genErr := generateEndpointID() + if genErr != nil { + return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} + } + if _, clash := seen[genID]; clash { + continue + } + if _, clash := existing[genID]; clash { + continue + } + useID = genID + seen[useID] = p.line + break + } + if useID == "" { + return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} + } + } + + loginPW := p.loginPW + if p.needGenerateLogin { + pw, genErr := generateLoginPassword() + if genErr != nil { + return nil, []csvLineError{{Line: p.line, Reason: "生成密码失败"}} + } + 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 +} + +func normalizeCSVHeader(cols []string) []string { + out := make([]string, len(cols)) + for i, c := range cols { + out[i] = strings.TrimSpace(strings.TrimPrefix(c, "\ufeff")) + } + return out +} diff --git a/internal/admin/endpoints_db.go b/internal/admin/endpoints_db.go new file mode 100644 index 0000000..2903306 --- /dev/null +++ b/internal/admin/endpoints_db.go @@ -0,0 +1,384 @@ +package admin + +import ( + "context" + "crypto/rand" + "database/sql" + "errors" + "strings" + "time" + "unicode/utf8" + + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +const idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789" + +type endpointInsert struct { + ID string + Name string + Remark string + Source string + LoginHash string + TalkHash sql.NullString + DefaultDelayMs int64 + CreatedAtMs int64 +} + +func generateEndpointID() (string, error) { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + return "", err + } + out := make([]byte, 8) + for i := range b { + out[i] = idAlphabet[int(b[i])%len(idAlphabet)] + } + return "e_" + string(out), nil +} + +func generateLoginPassword() (string, error) { + const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789" + for range 8 { + b := make([]byte, 20) + if _, err := rand.Read(b); err != nil { + return "", err + } + out := make([]byte, 20) + for i := range b { + out[i] = alphabet[int(b[i])%len(alphabet)] + } + pw := string(out) + if !strings.HasPrefix(pw, protocol.SessionTokenPrefix) && utf8.RuneCountInString(pw) >= protocol.MinLoginPasswordLen { + return pw, nil + } + } + return "", errors.New("admin: generate login password failed") +} + +func isUniqueConstraint(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "unique constraint") || strings.Contains(msg, "constraint failed") +} + +func (h *Handler) insertEndpoint(ctx context.Context, in endpointInsert) (string, error) { + nowMs := in.CreatedAtMs + if nowMs == 0 { + nowMs = time.Now().UnixMilli() + } + const maxAttempts = 8 + requestedID := in.ID + for attempt := 0; attempt < maxAttempts; attempt++ { + useID := requestedID + if useID == "" { + genID, err := generateEndpointID() + if err != nil { + return "", err + } + useID = genID + } + err := h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + _, execErr := tx.ExecContext(ctx, ` +INSERT INTO endpoints( + id, name, remark, source, login_hash, talk_hash, talk_version, + default_delay_ms, enabled, created_at +) VALUES (?, ?, ?, ?, ?, ?, 0, ?, 1, ?)`, + useID, in.Name, in.Remark, in.Source, in.LoginHash, in.TalkHash, in.DefaultDelayMs, nowMs, + ) + return execErr + }) + if err == nil { + return useID, nil + } + if isUniqueConstraint(err) { + if requestedID != "" { + return "", err + } + continue + } + return "", err + } + return "", errors.New("admin: generate endpoint id exhausted") +} + +func (h *Handler) insertEndpointsBatch(ctx context.Context, rows []endpointInsert) error { + nowMs := time.Now().UnixMilli() + return h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + for i := range rows { + in := rows[i] + created := in.CreatedAtMs + if created == 0 { + created = nowMs + } + if _, err := tx.ExecContext(ctx, ` +INSERT INTO endpoints( + id, name, remark, source, login_hash, talk_hash, talk_version, + default_delay_ms, enabled, created_at +) VALUES (?, ?, ?, ?, ?, ?, 0, ?, 1, ?)`, + in.ID, in.Name, in.Remark, in.Source, in.LoginHash, in.TalkHash, in.DefaultDelayMs, created, + ); err != nil { + return err + } + } + return nil + }) +} + +func onlineSQLExpr() string { + return `(online_since IS NOT NULL AND (offline_since IS NULL OR online_since > offline_since))` +} + +func (h *Handler) buildEndpointFilter(source string, online, enabled *bool, query string) (where string, args []any) { + conds := []string{"1=1"} + if source != "" { + conds = append(conds, "source = ?") + args = append(args, source) + } + if enabled != nil { + if *enabled { + conds = append(conds, "enabled = 1") + } else { + conds = append(conds, "enabled = 0") + } + } + if online != nil { + if *online { + conds = append(conds, onlineSQLExpr()) + } else { + conds = append(conds, "NOT "+onlineSQLExpr()) + } + } + if query != "" { + conds = append(conds, "(id LIKE ? COLLATE NOCASE OR instr(lower(name), lower(?)) > 0)") + args = append(args, query+"%", query) + } + return strings.Join(conds, " AND "), args +} + +func (h *Handler) countEndpoints(ctx context.Context, source string, online, enabled *bool, query string) (int, error) { + where, args := h.buildEndpointFilter(source, online, enabled, query) + var n int + err := h.db.Read.QueryRowContext(ctx, `SELECT COUNT(*) FROM endpoints WHERE `+where, args...).Scan(&n) + return n, err +} + +func (h *Handler) listEndpoints(ctx context.Context, source string, online, enabled *bool, query string, limit, offset int) ([]endpointRow, error) { + where, args := h.buildEndpointFilter(source, online, enabled, query) + args = append(args, limit, offset) + rows, err := h.db.Read.QueryContext(ctx, ` +SELECT id, name, remark, source, enabled, talk_hash, default_delay_ms, created_at, + online_since, offline_since, session_issued_at, session_used_at +FROM endpoints +WHERE `+where+` +ORDER BY id ASC +LIMIT ? OFFSET ?`, args...) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + out := make([]endpointRow, 0) + for rows.Next() { + var e endpointRow + var enabledInt int + if scanErr := rows.Scan( + &e.ID, &e.Name, &e.Remark, &e.Source, &enabledInt, &e.TalkHash, &e.DefaultDelayMs, &e.CreatedAtMs, + &e.OnlineSinceMs, &e.OfflineSinceMs, &e.SessionIssuedAt, &e.SessionUsedAt, + ); scanErr != nil { + return nil, scanErr + } + e.Enabled = enabledInt != 0 + out = append(out, e) + } + return out, rows.Err() +} + +func (h *Handler) getEndpoint(ctx context.Context, id string) (endpointRow, error) { + var e endpointRow + var enabledInt int + err := h.db.Read.QueryRowContext(ctx, ` +SELECT id, name, remark, source, enabled, talk_hash, default_delay_ms, created_at, + online_since, offline_since, session_issued_at, session_used_at +FROM endpoints WHERE id = ?`, id).Scan( + &e.ID, &e.Name, &e.Remark, &e.Source, &enabledInt, &e.TalkHash, &e.DefaultDelayMs, &e.CreatedAtMs, + &e.OnlineSinceMs, &e.OfflineSinceMs, &e.SessionIssuedAt, &e.SessionUsedAt, + ) + if err != nil { + return endpointRow{}, err + } + e.Enabled = enabledInt != 0 + return e, nil +} + +func (h *Handler) endpointExists(ctx context.Context, id string) (bool, error) { + var n int + err := h.db.Read.QueryRowContext(ctx, `SELECT COUNT(1) FROM endpoints WHERE id = ?`, id).Scan(&n) + return n > 0, err +} + +func (h *Handler) existingEndpointIDs(ctx context.Context, ids []string) (map[string]struct{}, error) { + out := make(map[string]struct{}) + if len(ids) == 0 { + return out, nil + } + // 分批 IN,避免超长 SQL;导入最多 1000。 + const chunk = 200 + for i := 0; i < len(ids); i += chunk { + end := i + chunk + if end > len(ids) { + end = len(ids) + } + part := ids[i:end] + placeholders := make([]string, len(part)) + args := make([]any, len(part)) + for j, id := range part { + placeholders[j] = "?" + args[j] = id + } + rows, err := h.db.Read.QueryContext(ctx, + `SELECT id FROM endpoints WHERE id IN (`+strings.Join(placeholders, ",")+`)`, args...) + if err != nil { + return nil, err + } + for rows.Next() { + var id string + if scanErr := rows.Scan(&id); scanErr != nil { + _ = rows.Close() + return nil, scanErr + } + out[id] = struct{}{} + } + err = rows.Err() + _ = rows.Close() + if err != nil { + return nil, err + } + } + return out, nil +} + +// patchEndpoint 更新字段。返回更新前是否 enabled,便于停用时踢线。 +func (h *Handler) patchEndpoint(ctx context.Context, id string, name, remark *string, delaySec *int64, enabled *bool) (wasEnabled bool, err error) { + err = h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + var enabledInt int + if qErr := tx.QueryRowContext(ctx, `SELECT enabled FROM endpoints WHERE id = ?`, id).Scan(&enabledInt); qErr != nil { + return qErr + } + wasEnabled = enabledInt != 0 + sets := make([]string, 0, 4) + args := make([]any, 0, 5) + if name != nil { + sets = append(sets, "name = ?") + args = append(args, *name) + } + if remark != nil { + sets = append(sets, "remark = ?") + args = append(args, *remark) + } + if delaySec != nil { + sets = append(sets, "default_delay_ms = ?") + args = append(args, (*delaySec)*1000) + } + if enabled != nil { + v := 0 + if *enabled { + v = 1 + } + sets = append(sets, "enabled = ?") + args = append(args, v) + if !*enabled { + sets = append(sets, "session_hash = NULL", "session_issued_at = NULL", "session_used_at = NULL") + } + } + args = append(args, id) + res, execErr := tx.ExecContext(ctx, `UPDATE endpoints SET `+strings.Join(sets, ", ")+` WHERE id = ?`, args...) + if execErr != nil { + return execErr + } + n, _ := res.RowsAffected() + if n == 0 { + return sql.ErrNoRows + } + return nil + }) + return wasEnabled, err +} + +func (h *Handler) setEndpointEnabled(ctx context.Context, id string, enabled bool) (found bool, err error) { + err = h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + v := 0 + if enabled { + v = 1 + } + query := `UPDATE endpoints SET enabled = ?` + args := []any{v} + if !enabled { + query += `, session_hash = NULL, session_issued_at = NULL, session_used_at = NULL` + } + query += ` WHERE id = ?` + args = append(args, id) + res, execErr := tx.ExecContext(ctx, query, args...) + if execErr != nil { + return execErr + } + n, _ := res.RowsAffected() + found = n > 0 + return nil + }) + return found, err +} + +// deleteEndpointBasic 删除端行并清掉与之相关的 talk_grants。 +// 作废消息、群主转让等完整级联留给 I5。 +func (h *Handler) deleteEndpointBasic(ctx context.Context, id string) (found bool, err error) { + err = h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + if _, e := tx.ExecContext(ctx, `DELETE FROM talk_grants WHERE sender_id = ? OR target_id = ?`, id, id); e != nil { + return e + } + res, e := tx.ExecContext(ctx, `DELETE FROM endpoints WHERE id = ?`, id) + if e != nil { + return e + } + n, _ := res.RowsAffected() + found = n > 0 + return nil + }) + return found, err +} + +func (h *Handler) resetLoginPassword(ctx context.Context, id, loginHash string) (bool, error) { + var found bool + err := h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + res, e := tx.ExecContext(ctx, ` +UPDATE endpoints +SET login_hash = ?, session_hash = NULL, session_issued_at = NULL, session_used_at = NULL +WHERE id = ?`, loginHash, id) + if e != nil { + return e + } + n, _ := res.RowsAffected() + found = n > 0 + return nil + }) + return found, err +} + +func (h *Handler) setTalkPassword(ctx context.Context, id string, talkHash sql.NullString) (bool, error) { + var found bool + err := h.db.Queue.Do(ctx, func(tx *sql.Tx) error { + res, e := tx.ExecContext(ctx, ` +UPDATE endpoints +SET talk_hash = ?, talk_version = talk_version + 1 +WHERE id = ?`, talkHash, id) + if e != nil { + return e + } + n, _ := res.RowsAffected() + found = n > 0 + return nil + }) + return found, err +} diff --git a/internal/admin/endpoints_test.go b/internal/admin/endpoints_test.go new file mode 100644 index 0000000..7a39d06 --- /dev/null +++ b/internal/admin/endpoints_test.go @@ -0,0 +1,364 @@ +package admin_test + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "io" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "path/filepath" + "strings" + "sync" + "testing" + + "git.asio.asia/nixevol/NixMsg/internal/admin" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +type kickRecorder struct { + mu sync.Mutex + Calls []string +} + +func (k *kickRecorder) Kick(_ context.Context, endpointID string) (bool, error) { + k.mu.Lock() + defer k.mu.Unlock() + k.Calls = append(k.Calls, endpointID) + return true, nil +} + +func (k *kickRecorder) count() int { + k.mu.Lock() + defer k.mu.Unlock() + return len(k.Calls) +} + +func setupEndpoints(t *testing.T) (*store.DB, *httptest.Server, *http.Client, *kickRecorder, *admin.MemoryLoginLocks) { + 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) + } + kick := &kickRecorder{} + locks := admin.NewMemoryLoginLocks() + h := admin.New(admin.Deps{ + DB: db, + Hash: hash, + Tokens: admin.NewRandomAPITokens(), + Locks: locks, + KickEndpoint: kick.Kick, + }) + 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, kick, locks +} + +func TestEndpointImportDuplicateRejectsAll(t *testing.T) { + db, srv, client, _, _ := setupEndpoints(t) + base := srv.URL + + csvBody := "" + + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + + "ep-a,甲,password1,,0,\n" + + "ep-b,乙,password2,,0,\n" + + "ep-a,丙,password3,,0,\n" + + 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) + } + 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 env struct { + OK bool `json:"ok"` + Error *struct { + Code string `json:"code"` + Message string `json:"message"` + } `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.Fatal(err) + } + if env.OK || env.Error == nil || env.Error.Code != "bad_request" { + t.Fatalf("env=%+v", env) + } + if env.Data == nil || len(env.Data.Errors) == 0 { + t.Fatalf("missing errors: %s", raw) + } + foundLine := false + for _, e := range env.Data.Errors { + if e.Line == 4 { + foundLine = true + break + } + } + if !foundLine { + t.Fatalf("want line 4 in errors: %+v", env.Data.Errors) + } + + var n int + if err := db.Read.QueryRow(`SELECT COUNT(*) FROM endpoints`).Scan(&n); err != nil { + t.Fatal(err) + } + if n != 0 { + t.Fatalf("want 0 endpoints created, got %d", n) + } +} + +func TestEndpointResetPasswordAppearsOnce(t *testing.T) { + db, srv, client, kick, _ := setupEndpoints(t) + base := srv.URL + + res := postJSON(t, client, base+"/api/admin/endpoints", + `{"id":"dev-1","name":"门口","login_password":"oldpass12"}`, + csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + + // 写入假会话令牌,确认重置会清空 + err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`UPDATE endpoints SET session_hash='abc', session_issued_at=1, session_used_at=1 WHERE id='dev-1'`) + return e + }) + if err != nil { + t.Fatal(err) + } + + res = postJSON(t, client, base+"/api/admin/endpoints/dev-1/reset-login-password", + `{}`, csrfHeaders()) + body, _ := io.ReadAll(res.Body) + _ = res.Body.Close() + if res.StatusCode != 200 { + t.Fatalf("reset status=%d body=%s", res.StatusCode, body) + } + var resetEnv struct { + OK bool `json:"ok"` + Data struct { + LoginPassword string `json:"login_password"` + } `json:"data"` + } + if err := json.Unmarshal(body, &resetEnv); err != nil { + t.Fatal(err) + } + if !resetEnv.OK || resetEnv.Data.LoginPassword == "" { + t.Fatalf("want one-time password, got %s", body) + } + pw := resetEnv.Data.LoginPassword + if strings.Count(string(body), pw) != 1 { + t.Fatalf("password should appear exactly once in response: %s", body) + } + if strings.HasPrefix(pw, "nst_") { + t.Fatalf("password must not start with nst_: %q", pw) + } + + var session sql.NullString + if err := db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id='dev-1'`).Scan(&session); err != nil { + t.Fatal(err) + } + if session.Valid { + t.Fatal("session_hash should be cleared") + } + if kick.count() < 1 { + t.Fatal("expected kick hook after reset") + } + + // 详情中不得再出现明文密码 + res = doReq(t, client, http.MethodGet, base+"/api/admin/endpoints/dev-1", "", nil) + detailBody, _ := io.ReadAll(res.Body) + _ = res.Body.Close() + if strings.Contains(string(detailBody), pw) { + t.Fatalf("password leaked in detail: %s", detailBody) + } +} + +func TestEndpointMutatingRequiresCSRF(t *testing.T) { + _, srv, client, _, _ := setupEndpoints(t) + base := srv.URL + + res := postJSON(t, client, base+"/api/admin/endpoints", + `{"id":"no-csrf","name":"x","login_password":"password1"}`, + nil) // 无 CSRF + env := decodeEnv(t, res) + if res.StatusCode != http.StatusForbidden || env.Error == nil || env.Error.Code != "forbidden" { + t.Fatalf("want 403 forbidden got %d %+v", res.StatusCode, env) + } +} + +func TestEndpointDisableKickAndUnlock(t *testing.T) { + db, srv, client, kick, locks := setupEndpoints(t) + base := srv.URL + + res := postJSON(t, client, base+"/api/admin/endpoints", + `{"id":"lock-1","name":"锁","login_password":"password1"}`, + csrfHeaders()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + + res = postJSON(t, client, base+"/api/admin/endpoints/batch", + `{"ids":["lock-1"],"action":"disable"}`, csrfHeaders()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("disable: %d %+v", res.StatusCode, env) + } + var enabled int + if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='lock-1'`).Scan(&enabled); err != nil { + t.Fatal(err) + } + if enabled != 0 { + t.Fatalf("want enabled=0 got %d", enabled) + } + if kick.count() < 1 { + t.Fatal("disable should kick") + } + + for i := 0; i < 50; i++ { + locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"}) + } + if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"}); !locked { + t.Fatal("expected endpoint locked before unlock") + } + + res = postJSON(t, client, base+"/api/admin/endpoints/lock-1/unlock", `{}`, csrfHeaders()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("unlock: %d %+v", res.StatusCode, env) + } + if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"}); locked { + t.Fatal("expected unlocked") + } + + res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password", + `{"talk_password":"talk"}`, + map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"}) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("talk-password: %d %+v", res.StatusCode, env) + } + var talkSet struct { + TalkPasswordSet bool `json:"talk_password_set"` + } + _ = json.Unmarshal(env.Data, &talkSet) + if !talkSet.TalkPasswordSet { + t.Fatal("want talk_password_set true") + } + + var ver int + if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil { + t.Fatal(err) + } + if ver != 1 { + t.Fatalf("talk_version want 1 got %d", ver) + } + + res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password", + `{"talk_password":""}`, + map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"}) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("clear talk: %d %+v", res.StatusCode, env) + } + if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil { + t.Fatal(err) + } + if ver != 2 { + t.Fatalf("talk_version want 2 got %d", ver) + } +} + +func TestEndpointImportBOMAndCreate(t *testing.T) { + db, srv, client, _, _ := setupEndpoints(t) + base := srv.URL + + var buf bytes.Buffer + buf.Write([]byte{0xEF, 0xBB, 0xBF}) + buf.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") + buf.WriteString("bom-1,门,,secret,5,备注\n") + + req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/import", &buf) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "text/csv; charset=utf-8") + req.Header.Set("X-Nixmsg-Request", "1") + res, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("import: %d %+v", res.StatusCode, env) + } + var data struct { + Items []struct { + ID string `json:"id"` + LoginPassword string `json:"login_password"` + Name string `json:"name"` + } `json:"items"` + } + if err := json.Unmarshal(env.Data, &data); err != nil { + t.Fatal(err) + } + if len(data.Items) != 1 || data.Items[0].ID != "bom-1" || data.Items[0].LoginPassword == "" { + t.Fatalf("items=%+v", data.Items) + } + + var n int + if err := db.Read.QueryRow(`SELECT COUNT(*) FROM endpoints WHERE id='bom-1'`).Scan(&n); err != nil { + t.Fatal(err) + } + if n != 1 { + t.Fatalf("want 1 row got %d", n) + } + + res = doReq(t, client, http.MethodGet, base+"/api/admin/endpoints?source=admin", "", nil) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("list: %d %+v", res.StatusCode, env) + } +} + +func csrfHeaders() map[string]string { + return map[string]string{"X-Nixmsg-Request": "1"} +} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index f66d13e..e49ed71 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -36,6 +36,8 @@ type Deps struct { SessionTTL time.Duration // SecureCookies 为 true 时 Cookie 始终带 Secure;否则按请求是否 HTTPS 决定。 SecureCookies bool + // KickEndpoint 踢下线钩子(只断开连接);nil 时踢线为 no-op。 + KickEndpoint EndpointKickFunc } // Handler 是可挂载的管理接口(路由前缀 /api/admin/)。 @@ -48,6 +50,7 @@ type Handler struct { trusted []*net.IPNet ttl time.Duration forceSec bool + kick EndpointKickFunc mux *http.ServeMux @@ -76,6 +79,7 @@ func New(d Deps) *Handler { trusted: d.TrustedProxies, ttl: ttl, forceSec: d.SecureCookies, + kick: d.KickEndpoint, mux: http.NewServeMux(), lastUsed: make(map[string]time.Time), } @@ -103,7 +107,20 @@ func (h *Handler) routes() { h.mux.Handle("PATCH /api/admin/tokens/{id}", h.auth(h.handleTokenPatch)) h.mux.Handle("DELETE /api/admin/tokens/{id}", h.auth(h.handleTokenDelete)) - // 其余管理路由:鉴权生效,业务暂 501 + // A2 端管理 + h.mux.Handle("GET /api/admin/endpoints", h.auth(h.handleEndpointList)) + h.mux.Handle("POST /api/admin/endpoints", h.auth(h.handleEndpointCreate)) + h.mux.Handle("POST /api/admin/endpoints/import", h.auth(h.handleEndpointImport)) + h.mux.Handle("POST /api/admin/endpoints/batch", h.auth(h.handleEndpointBatch)) + h.mux.Handle("GET /api/admin/endpoints/{id}", h.auth(h.handleEndpointGet)) + h.mux.Handle("PATCH /api/admin/endpoints/{id}", h.auth(h.handleEndpointPatch)) + h.mux.Handle("DELETE /api/admin/endpoints/{id}", h.auth(h.handleEndpointDelete)) + h.mux.Handle("POST /api/admin/endpoints/{id}/kick", h.auth(h.handleEndpointKick)) + h.mux.Handle("POST /api/admin/endpoints/{id}/reset-login-password", h.auth(h.handleEndpointResetLoginPassword)) + h.mux.Handle("PUT /api/admin/endpoints/{id}/talk-password", h.auth(h.handleEndpointTalkPassword)) + h.mux.Handle("POST /api/admin/endpoints/{id}/unlock", h.auth(h.handleEndpointUnlock)) + + // 其余管理路由:鉴权生效,业务暂 501(A3) for _, p := range stubRoutes { h.mux.Handle(p, h.auth(h.handleNotImplemented)) } @@ -111,17 +128,6 @@ func (h *Handler) routes() { var stubRoutes = []string{ "GET /api/admin/overview", - "GET /api/admin/endpoints", - "POST /api/admin/endpoints", - "POST /api/admin/endpoints/import", - "POST /api/admin/endpoints/batch", - "GET /api/admin/endpoints/{id}", - "PATCH /api/admin/endpoints/{id}", - "DELETE /api/admin/endpoints/{id}", - "POST /api/admin/endpoints/{id}/kick", - "POST /api/admin/endpoints/{id}/reset-login-password", - "PUT /api/admin/endpoints/{id}/talk-password", - "POST /api/admin/endpoints/{id}/unlock", "GET /api/admin/registration", "PUT /api/admin/registration", "GET /api/admin/groups",