Files

224 lines
4.7 KiB
Go

package settings
import (
"context"
"encoding/json"
"strings"
"github.com/jackc/pgx/v5/pgxpool"
"amnezia-share/internal/panel"
)
const (
KeyPanelURL = "amnezia_panel_url"
KeyAPIToken = "amnezia_api_token"
KeyServerLabelsJSON = "amnezia_server_labels_json"
KeyTrafficNotices = "amnezia_share_server_traffic_notices_json"
KeyProtocolsJSON = "amnezia_server_protocols_json"
KeyFlagsJSON = "amnezia_server_flags_json"
KeyDisabledServers = "amnezia_disabled_servers_json"
KeySpeedsJSON = "amnezia_server_speeds_json"
KeyMaintenance = "share_maintenance_mode"
KeyPocketURL = "pocket_id_url"
KeyPocketClientID = "pocket_id_client_id"
KeyPocketSecret = "pocket_id_client_secret"
KeyMemberDays = "member_default_days"
KeyMemberMaxUses = "member_default_max_uses"
)
type Store struct {
Pool *pgxpool.Pool
EnvPanelURL string
EnvToken string
EnvLabelsJSON string
}
func (s *Store) Get(ctx context.Context, key string) (string, error) {
var v *string
err := s.Pool.QueryRow(ctx, `SELECT svalue FROM site_settings WHERE skey=$1`, key).Scan(&v)
if err != nil {
return "", nil
}
if v == nil {
return "", nil
}
return *v, nil
}
func (s *Store) Set(ctx context.Context, key, value string) error {
_, err := s.Pool.Exec(ctx, `
INSERT INTO site_settings (skey, svalue, updated_at) VALUES ($1,$2,NOW())
ON CONFLICT (skey) DO UPDATE SET svalue=EXCLUDED.svalue, updated_at=NOW()`, key, value)
return err
}
func (s *Store) GetMany(ctx context.Context, keys ...string) (map[string]string, error) {
out := map[string]string{}
for _, k := range keys {
out[k] = ""
}
rows, err := s.Pool.Query(ctx, `SELECT skey, svalue FROM site_settings WHERE skey = ANY($1)`, keys)
if err != nil {
return out, nil
}
defer rows.Close()
for rows.Next() {
var k string
var v *string
if err := rows.Scan(&k, &v); err != nil {
continue
}
if v != nil {
out[k] = *v
}
}
return out, nil
}
func (s *Store) PanelURL(ctx context.Context) string {
v, _ := s.Get(ctx, KeyPanelURL)
v = strings.TrimSpace(v)
if v != "" {
return v
}
return strings.TrimSpace(s.EnvPanelURL)
}
func (s *Store) PanelToken(ctx context.Context) string {
v, _ := s.Get(ctx, KeyAPIToken)
v = panel.NormalizeToken(v)
if v != "" {
return v
}
return panel.NormalizeToken(s.EnvToken)
}
func (s *Store) Maintenance(ctx context.Context) bool {
v, _ := s.Get(ctx, KeyMaintenance)
return v == "1"
}
func (s *Store) ServerLabels(ctx context.Context) map[int]string {
out := map[int]string{}
merge := func(raw string) {
raw = strings.TrimSpace(raw)
if raw == "" {
return
}
var m map[string]any
if json.Unmarshal([]byte(raw), &m) != nil {
return
}
for k, v := range m {
var id int
fmtSscanf(k, &id)
switch t := v.(type) {
case string:
if strings.TrimSpace(t) != "" {
out[id] = strings.TrimSpace(t)
}
}
}
}
merge(s.EnvLabelsJSON)
dbJSON, _ := s.Get(ctx, KeyServerLabelsJSON)
merge(dbJSON)
rows, err := s.Pool.Query(ctx, `SELECT panel_server_id, title FROM panel_server_labels`)
if err == nil {
defer rows.Close()
for rows.Next() {
var id int
var title string
if rows.Scan(&id, &title) == nil && strings.TrimSpace(title) != "" {
out[id] = strings.TrimSpace(title)
}
}
}
return out
}
func (s *Store) JSONMapString(ctx context.Context, key string) map[int]string {
out := map[int]string{}
raw, _ := s.Get(ctx, key)
raw = strings.TrimSpace(raw)
if raw == "" {
return out
}
var m map[string]any
if json.Unmarshal([]byte(raw), &m) != nil {
return out
}
for k, v := range m {
var id int
fmtSscanf(k, &id)
if str, ok := v.(string); ok {
out[id] = str
}
}
return out
}
func (s *Store) JSONMapStringSlice(ctx context.Context, key string) map[int][]string {
out := map[int][]string{}
raw, _ := s.Get(ctx, key)
raw = strings.TrimSpace(raw)
if raw == "" {
return out
}
var m map[string]any
if json.Unmarshal([]byte(raw), &m) != nil {
return out
}
for k, v := range m {
var id int
fmtSscanf(k, &id)
arr, ok := v.([]any)
if !ok {
continue
}
var list []string
for _, x := range arr {
if str, ok := x.(string); ok {
list = append(list, str)
}
}
out[id] = list
}
return out
}
func (s *Store) DisabledServers(ctx context.Context) map[int]bool {
out := map[int]bool{}
raw, _ := s.Get(ctx, KeyDisabledServers)
raw = strings.TrimSpace(raw)
if raw == "" {
return out
}
var arr []any
if json.Unmarshal([]byte(raw), &arr) != nil {
return out
}
for _, x := range arr {
switch t := x.(type) {
case float64:
out[int(t)] = true
case int:
out[t] = true
}
}
return out
}
func fmtSscanf(s string, id *int) {
n := 0
for _, r := range s {
if r < '0' || r > '9' {
break
}
n = n*10 + int(r-'0')
}
*id = n
}