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 }