feat: 实现管理后台端管理接口(A2)
This commit is contained in:
+31
-4
@@ -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`**
|
||||
|
||||
@@ -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},
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"}
|
||||
}
|
||||
+18
-12
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user