224 lines
4.7 KiB
Go
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
|
|
}
|