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) passwordResetKick(ctx context.Context, id string) (bool, error) { if h.resetKick != nil { return h.resetKick(ctx, id) } return h.kickEndpoint(ctx, id) } func (h *Handler) afterDisableKick(ctx context.Context, id string) { if h.disableKick != nil { _, _ = h.disableKick(ctx, id) return } if h.identity == nil { _, _ = h.kickEndpoint(ctx, id) } } func (h *Handler) afterDeleteKick(ctx context.Context, id string) { if h.deleteKick != nil { _, _ = h.deleteKick(ctx, id) return } if h.identity == nil { _, _ = h.kickEndpoint(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, h.maxScheduleSeconds()); 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 { if msg := validateDelaySeconds(*req.DefaultDelaySeconds, h.maxScheduleSeconds()); msg != "" { h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", msg) return } } hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil var wasEnabled bool var err error // 注入 Identity 时启停只走 identity,避免先写 enabled 再级联失败造成半生效。 patchEnabled := req.Enabled if h.identity != nil { patchEnabled = nil } if hasMeta || (req.Enabled != nil && h.identity == nil) { wasEnabled, err = h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, patchEnabled) 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 && h.identity != nil { var found bool found, err = h.setEndpointEnabled(r.Context(), id, *req.Enabled) if err != nil { h.audit(actorString(p), "endpoint_patch", id, "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } if !found { h.audit(actorString(p), "endpoint_patch", id, "not_found", ip) httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") return } if !*req.Enabled { h.afterDisableKick(r.Context(), id) } } else 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.afterDeleteKick(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.afterDisableKick(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.afterDeleteKick(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.passwordResetKick(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, maxDelaySec 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 msg := validateDelaySeconds(delaySec, maxDelaySec); msg != "" { return msg } return "" } func validateDelaySeconds(delaySec, maxDelaySec int64) string { if delaySec < 0 { return "默认延迟无效" } if maxDelaySec > 0 && delaySec > maxDelaySec { return "默认延迟无效" } return "" } func (h *Handler) maxScheduleSeconds() int64 { return int64(h.cfg.Limits.MaxScheduleSeconds) } func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{ OK: false, Error: &httpx.ErrorBody{ Code: "bad_request", Message: "CSV 校验失败", }, Data: map[string]any{"errors": errs}, }) } func writeCSVConflictError(w http.ResponseWriter, errs []csvLineError) { httpx.WriteJSON(w, http.StatusConflict, httpx.Envelope{ OK: false, Error: &httpx.ErrorBody{ Code: "id_taken", Message: "编号已占用", }, Data: map[string]any{"errors": errs}, }) }