Files

473 lines
14 KiB
Go

package admin
import (
"context"
"database/sql"
"errors"
"net/http"
"strconv"
"strings"
"git.asio.asia/nixevol/NixMsg/internal/app/group"
"git.asio.asia/nixevol/NixMsg/internal/httpx"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
func (h *Handler) requireGroups(w http.ResponseWriter) bool {
if h.groups == nil {
httpx.WriteError(w, http.StatusServiceUnavailable, "busy", "群服务未注入")
return false
}
return true
}
func (h *Handler) handleGroupList(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query()
limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor"))
if !ok {
return
}
query := strings.TrimSpace(q.Get("query"))
where := "1=1"
args := make([]any, 0, 4)
if query != "" {
where += " AND LOWER(g.name) LIKE ?"
args = append(args, "%"+strings.ToLower(query)+"%")
}
var total int
countSQL := `SELECT COUNT(*) FROM groups g WHERE ` + where
if err := h.db.Read.QueryRowContext(r.Context(), countSQL, args...).Scan(&total); err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
listArgs := append(append([]any{}, args...), limit, offset)
rows, err := h.db.Read.QueryContext(r.Context(), `
SELECT g.id, g.name, g.owner_id, g.created_at,
(SELECT COUNT(*) FROM group_members gm WHERE gm.group_id = g.id) AS cnt
FROM groups g
WHERE `+where+`
ORDER BY g.id ASC
LIMIT ? OFFSET ?`, listArgs...)
if err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
defer func() { _ = rows.Close() }()
items := make([]map[string]any, 0)
for rows.Next() {
var id, name, owner string
var created int64
var cnt int
if scanErr := rows.Scan(&id, &name, &owner, &created, &cnt); scanErr != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
items = append(items, map[string]any{
"id": id,
"name": name,
"owner_id": owner,
"member_count": cnt,
"created_at_ms": created,
})
}
if err := rows.Err(); err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
next := ""
if offset+len(items) < total {
next = strconv.Itoa(offset + len(items))
}
httpx.WriteOK(w, map[string]any{"items": items, "next_cursor": next, "total": total})
}
func (h *Handler) handleGroupCreate(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
var req struct {
ID string `json:"id"`
Name string `json:"name"`
OwnerID string `json:"owner_id"`
MemberIDs []string `json:"member_ids"`
}
if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "group_create", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
return
}
if strings.TrimSpace(req.ID) != "" {
// AdminCreate 不接受自定义 id;与契约「留空则生成」一致时忽略非空会误导,故拒绝。
h.audit(actorString(p), "group_create", req.ID, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "后台创建群请留空 id,由服务器生成")
return
}
if !protocol.ValidName(req.Name) || req.Name == "" {
h.audit(actorString(p), "group_create", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "名称不合法")
return
}
if req.OwnerID == "" {
h.audit(actorString(p), "group_create", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "缺少 owner_id")
return
}
res, err := h.groups.AdminCreate(r.Context(), req.Name, req.OwnerID, req.MemberIDs)
if err != nil {
h.writeGroupErr(w, p, "group_create", req.OwnerID, ip, err)
return
}
h.audit(actorString(p), "group_create", res.ID, "ok", ip)
failed := res.Failed
if failed == nil {
failed = []group.MemberFail{}
}
httpx.WriteOK(w, map[string]any{
"id": res.ID,
"name": res.Name,
"owner_id": res.OwnerID,
"failed": failed,
})
}
func (h *Handler) handleGroupGet(w http.ResponseWriter, r *http.Request) {
id := r.PathValue("id")
q := r.URL.Query()
limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor"))
if !ok {
return
}
var name, owner string
var created int64
err := h.db.Read.QueryRowContext(r.Context(),
`SELECT name, owner_id, created_at FROM groups WHERE id = ?`, id,
).Scan(&name, &owner, &created)
if isNoRows(err) {
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
}
if err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
rows, err := h.db.Read.QueryContext(r.Context(), `
SELECT gm.endpoint_id, COALESCE(e.name, ''), gm.joined_at,
e.online_since, e.offline_since
FROM group_members gm
LEFT JOIN endpoints e ON e.id = gm.endpoint_id
WHERE gm.group_id = ?
ORDER BY gm.endpoint_id ASC
LIMIT ? OFFSET ?`, id, limit, offset)
if err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
defer func() { _ = rows.Close() }()
members := make([]map[string]any, 0)
for rows.Next() {
var eid, ename string
var joined int64
var onlineSince, offlineSince sql.NullInt64
if scanErr := rows.Scan(&eid, &ename, &joined, &onlineSince, &offlineSince); scanErr != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
online, _, _ := presenceFields(onlineSince, offlineSince)
members = append(members, map[string]any{
"id": eid,
"name": ename,
"online": online,
"joined_at_ms": joined,
})
}
if err := rows.Err(); err != nil {
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
var memberTotal int
_ = h.db.Read.QueryRowContext(r.Context(),
`SELECT COUNT(*) FROM group_members WHERE group_id = ?`, id).Scan(&memberTotal)
next := ""
if offset+len(members) < memberTotal {
next = strconv.Itoa(offset + len(members))
}
httpx.WriteOK(w, map[string]any{
"id": id,
"name": name,
"owner_id": owner,
"created_at_ms": created,
"members": members,
"next_cursor": next,
})
}
func (h *Handler) handleGroupRename(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
id := r.PathValue("id")
var req struct {
Name string `json:"name"`
}
if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "group_rename", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
return
}
owner, err := h.groupOwner(r.Context(), id)
if isNoRows(err) {
h.audit(actorString(p), "group_rename", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
}
if err != nil {
h.audit(actorString(p), "group_rename", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
err = h.groups.Rename(r.Context(), owner, &protocol.GroupRename{
V: protocol.Version, Type: protocol.TypeGroupRename, RID: "admin",
GroupID: id, Name: req.Name,
})
if err != nil {
h.writeGroupErr(w, p, "group_rename", id, ip, err)
return
}
summary, err := h.groupSummary(r.Context(), id)
if err != nil {
h.audit(actorString(p), "group_rename", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
h.audit(actorString(p), "group_rename", id, "ok", ip)
httpx.WriteOK(w, summary)
}
func (h *Handler) handleGroupDissolve(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
id := r.PathValue("id")
owner, err := h.groupOwner(r.Context(), id)
if isNoRows(err) {
h.audit(actorString(p), "group_dissolve", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
}
if err != nil {
h.audit(actorString(p), "group_dissolve", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
err = h.groups.Dissolve(r.Context(), owner, &protocol.GroupDissolve{
V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "admin", GroupID: id,
})
if err != nil {
h.writeGroupErr(w, p, "group_dissolve", id, ip, err)
return
}
h.audit(actorString(p), "group_dissolve", id, "ok", ip)
httpx.WriteOK(w, map[string]any{})
}
func (h *Handler) handleGroupAddMembers(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
id := r.PathValue("id")
var req struct {
MemberIDs []string `json:"member_ids"`
}
if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "group_add_members", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
return
}
if _, err := h.groupOwner(r.Context(), id); isNoRows(err) {
h.audit(actorString(p), "group_add_members", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
} else if err != nil {
h.audit(actorString(p), "group_add_members", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
res, err := h.groups.AdminAddMembers(r.Context(), id, req.MemberIDs)
if err != nil {
h.writeGroupErr(w, p, "group_add_members", id, ip, err)
return
}
failed := res.Failed
if failed == nil {
failed = []group.MemberFail{}
}
h.audit(actorString(p), "group_add_members", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"failed": failed})
}
func (h *Handler) handleGroupRemoveMember(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
id := r.PathValue("id")
endpointID := r.PathValue("endpointId")
owner, err := h.groupOwner(r.Context(), id)
if isNoRows(err) {
h.audit(actorString(p), "group_remove_member", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
}
if err != nil {
h.audit(actorString(p), "group_remove_member", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
err = h.groups.Remove(r.Context(), owner, &protocol.GroupRemove{
V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "admin",
GroupID: id, EndpointID: endpointID,
})
if err != nil {
h.writeGroupErr(w, p, "group_remove_member", id, ip, err)
return
}
h.audit(actorString(p), "group_remove_member", id+"/"+endpointID, "ok", ip)
httpx.WriteOK(w, map[string]any{})
}
func (h *Handler) handleGroupTransfer(w http.ResponseWriter, r *http.Request) {
if !h.requireGroups(w) {
return
}
p, _ := principalFrom(r.Context())
ip := httpx.ClientIP(r, h.trusted)
id := r.PathValue("id")
var req struct {
EndpointID string `json:"endpoint_id"`
}
if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "group_transfer", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
return
}
owner, err := h.groupOwner(r.Context(), id)
if isNoRows(err) {
h.audit(actorString(p), "group_transfer", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在")
return
}
if err != nil {
h.audit(actorString(p), "group_transfer", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
err = h.groups.Transfer(r.Context(), owner, &protocol.GroupTransfer{
V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "admin",
GroupID: id, EndpointID: req.EndpointID,
})
if err != nil {
h.writeGroupErr(w, p, "group_transfer", id, ip, err)
return
}
h.audit(actorString(p), "group_transfer", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"owner_id": req.EndpointID})
}
func (h *Handler) groupOwner(ctx context.Context, groupID string) (string, error) {
var owner string
err := h.db.Read.QueryRowContext(ctx, `SELECT owner_id FROM groups WHERE id = ?`, groupID).Scan(&owner)
return owner, err
}
func (h *Handler) groupSummary(ctx context.Context, groupID string) (map[string]any, error) {
var name, owner string
var created int64
var cnt int
err := h.db.Read.QueryRowContext(ctx, `
SELECT g.name, g.owner_id, g.created_at,
(SELECT COUNT(*) FROM group_members gm WHERE gm.group_id = g.id)
FROM groups g WHERE g.id = ?`, groupID).Scan(&name, &owner, &created, &cnt)
if err != nil {
return nil, err
}
return map[string]any{
"id": groupID,
"name": name,
"owner_id": owner,
"member_count": cnt,
"created_at_ms": created,
}, nil
}
func (h *Handler) writeGroupErr(w http.ResponseWriter, p principal, action, object, ip string, err error) {
var pe *protocol.Error
if errors.As(err, &pe) {
status := http.StatusBadRequest
switch pe.Code {
case protocol.CodeNotFound, protocol.CodeInvalidTarget:
status = http.StatusNotFound
case protocol.CodeForbidden, protocol.CodeNotMember, protocol.CodeOwnerCannotLeave:
status = http.StatusForbidden
case protocol.CodeIDTaken:
status = http.StatusConflict
case protocol.CodeBusy:
status = http.StatusServiceUnavailable
}
h.audit(actorString(p), action, object, pe.Code, ip)
httpx.WriteError(w, status, pe.Code, pe.Message)
return
}
h.audit(actorString(p), action, object, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
}
func parsePage(w http.ResponseWriter, limitStr, cursorStr string) (limit, offset int, ok bool) {
limit = defaultListLimit
if limitStr != "" {
n, err := strconv.Atoi(limitStr)
if err != nil || n < 1 {
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "limit 无效")
return 0, 0, false
}
limit = n
}
if limit > maxListLimit {
limit = maxListLimit
}
if cursorStr != "" {
n, err := strconv.Atoi(cursorStr)
if err != nil || n < 0 {
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "cursor 无效")
return 0, 0, false
}
offset = n
}
return limit, offset, true
}