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 }