Files
NixMsg/internal/admin/endpoints_db.go
T

414 lines
11 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) {
if h.identity != nil {
var opErr error
if enabled {
opErr = h.identity.Enable(ctx, id)
} else {
opErr = h.identity.Disable(ctx, id)
}
if opErr != nil {
if isEndpointNotFound(opErr) {
return false, nil
}
return false, opErr
}
return true, nil
}
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 删除端;注入 Identity 时走 I5 完整级联。
func (h *Handler) deleteEndpointBasic(ctx context.Context, id string) (found bool, err error) {
if h.identity != nil {
opErr := h.identity.Delete(ctx, id)
if opErr != nil {
if isEndpointNotFound(opErr) {
return false, nil
}
return false, opErr
}
return true, nil
}
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 isEndpointNotFound(err error) bool {
var pe *protocol.Error
return errors.As(err, &pe) && pe.Code == protocol.CodeNotFound
}
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
}