feat: 实现管理后台端管理接口(A2)
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user