385 lines
10 KiB
Go
385 lines
10 KiB
Go
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
|
|
}
|