mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
913 lines
30 KiB
Go
913 lines
30 KiB
Go
package service
|
||
|
||
import (
|
||
"context"
|
||
"crypto/rand"
|
||
"encoding/hex"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"log/slog"
|
||
"math"
|
||
"strconv"
|
||
"strings"
|
||
)
|
||
|
||
// IsRegistrationEnabled 检查是否开放注册
|
||
func (s *SettingService) IsRegistrationEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyRegistrationEnabled)
|
||
if err != nil {
|
||
// 安全默认:如果设置不存在或查询出错,默认关闭注册
|
||
return false
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// IsEmailVerifyEnabled 检查是否开启邮件验证
|
||
func (s *SettingService) IsEmailVerifyEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyEmailVerifyEnabled)
|
||
if err != nil {
|
||
return false
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetRegistrationEmailSuffixWhitelist returns normalized registration email suffix whitelist.
|
||
func (s *SettingService) GetRegistrationEmailSuffixWhitelist(ctx context.Context) []string {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyRegistrationEmailSuffixWhitelist)
|
||
if err != nil {
|
||
return []string{}
|
||
}
|
||
return ParseRegistrationEmailSuffixWhitelist(value)
|
||
}
|
||
|
||
// IsPromoCodeEnabled 检查是否启用优惠码功能
|
||
func (s *SettingService) IsPromoCodeEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyPromoCodeEnabled)
|
||
if err != nil {
|
||
return true // 默认启用
|
||
}
|
||
return value != "false"
|
||
}
|
||
|
||
// IsInvitationCodeEnabled 检查是否启用邀请码注册功能
|
||
func (s *SettingService) IsInvitationCodeEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyInvitationCodeEnabled)
|
||
if err != nil {
|
||
return false // 默认关闭
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetCustomMenuItemsRaw returns the raw JSON string of custom_menu_items setting.
|
||
func (s *SettingService) GetCustomMenuItemsRaw(ctx context.Context) string {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyCustomMenuItems)
|
||
if err != nil {
|
||
return "[]"
|
||
}
|
||
return value
|
||
}
|
||
|
||
// IsAffiliateEnabled 检查是否启用邀请返利功能(总开关)
|
||
func (s *SettingService) IsAffiliateEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyAffiliateEnabled)
|
||
if err != nil {
|
||
return false // 默认关闭
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetAffiliateRebateRatePercent 读取并 clamp 全局返利比例。
|
||
// 解析失败、缺失或越界都回退到 AffiliateRebateRateDefault — 该比例从不抛错,
|
||
// 调用方只关心一个可用的数值。
|
||
func (s *SettingService) GetAffiliateRebateRatePercent(ctx context.Context) float64 {
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyAffiliateRebateRate)
|
||
if err != nil {
|
||
return AffiliateRebateRateDefault
|
||
}
|
||
rate, err := strconv.ParseFloat(strings.TrimSpace(raw), 64)
|
||
if err != nil || math.IsNaN(rate) || math.IsInf(rate, 0) {
|
||
return AffiliateRebateRateDefault
|
||
}
|
||
return clampAffiliateRebateRate(rate)
|
||
}
|
||
|
||
// GetAffiliateRebateFreezeHours 返回返利冻结期(小时)。
|
||
// 返回 0 表示不冻结(向后兼容)。
|
||
func (s *SettingService) GetAffiliateRebateFreezeHours(ctx context.Context) int {
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyAffiliateRebateFreezeHours)
|
||
if err != nil {
|
||
return AffiliateRebateFreezeHoursDefault
|
||
}
|
||
hours, err := strconv.Atoi(strings.TrimSpace(raw))
|
||
if err != nil || hours < 0 {
|
||
return AffiliateRebateFreezeHoursDefault
|
||
}
|
||
if hours > AffiliateRebateFreezeHoursMax {
|
||
return AffiliateRebateFreezeHoursMax
|
||
}
|
||
return hours
|
||
}
|
||
|
||
// GetAffiliateRebateDurationDays 返回返利有效期(天)。
|
||
// 返回 0 表示永久有效。
|
||
func (s *SettingService) GetAffiliateRebateDurationDays(ctx context.Context) int {
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyAffiliateRebateDurationDays)
|
||
if err != nil {
|
||
return AffiliateRebateDurationDaysDefault
|
||
}
|
||
days, err := strconv.Atoi(strings.TrimSpace(raw))
|
||
if err != nil || days < 0 {
|
||
return AffiliateRebateDurationDaysDefault
|
||
}
|
||
if days > AffiliateRebateDurationDaysMax {
|
||
return AffiliateRebateDurationDaysMax
|
||
}
|
||
return days
|
||
}
|
||
|
||
// GetAffiliateRebatePerInviteeCap 返回单人返利上限。
|
||
// 返回 0 表示无上限。
|
||
func (s *SettingService) GetAffiliateRebatePerInviteeCap(ctx context.Context) float64 {
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyAffiliateRebatePerInviteeCap)
|
||
if err != nil {
|
||
return AffiliateRebatePerInviteeCapDefault
|
||
}
|
||
cap, err := strconv.ParseFloat(strings.TrimSpace(raw), 64)
|
||
if err != nil || cap < 0 || math.IsNaN(cap) || math.IsInf(cap, 0) {
|
||
return AffiliateRebatePerInviteeCapDefault
|
||
}
|
||
return cap
|
||
}
|
||
|
||
// IsPasswordResetEnabled 检查是否启用密码重置功能
|
||
// 要求:必须同时开启邮件验证
|
||
func (s *SettingService) IsPasswordResetEnabled(ctx context.Context) bool {
|
||
// Password reset requires email verification to be enabled
|
||
if !s.IsEmailVerifyEnabled(ctx) {
|
||
return false
|
||
}
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyPasswordResetEnabled)
|
||
if err != nil {
|
||
return false // 默认关闭
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// IsTotpEnabled 检查是否启用 TOTP 双因素认证功能
|
||
func (s *SettingService) IsTotpEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyTotpEnabled)
|
||
if err != nil {
|
||
return false // 默认关闭
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// IsTotpEncryptionKeyConfigured 检查 TOTP 加密密钥是否已手动配置
|
||
// 只有手动配置了密钥才允许在管理后台启用 TOTP 功能
|
||
func (s *SettingService) IsTotpEncryptionKeyConfigured() bool {
|
||
return s.cfg.Totp.EncryptionKeyConfigured
|
||
}
|
||
|
||
// GetSiteName 获取网站名称
|
||
func (s *SettingService) GetSiteName(ctx context.Context) string {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeySiteName)
|
||
if err != nil || value == "" {
|
||
return "Sub2API"
|
||
}
|
||
return value
|
||
}
|
||
|
||
// GetDefaultConcurrency 获取默认并发量
|
||
func (s *SettingService) GetDefaultConcurrency(ctx context.Context) int {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyDefaultConcurrency)
|
||
if err != nil {
|
||
return s.cfg.Default.UserConcurrency
|
||
}
|
||
if v, err := strconv.Atoi(value); err == nil && v > 0 {
|
||
return v
|
||
}
|
||
return s.cfg.Default.UserConcurrency
|
||
}
|
||
|
||
// GetDefaultBalance 获取默认余额
|
||
func (s *SettingService) GetDefaultBalance(ctx context.Context) float64 {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyDefaultBalance)
|
||
if err != nil {
|
||
return s.cfg.Default.UserBalance
|
||
}
|
||
if v, err := strconv.ParseFloat(value, 64); err == nil && v >= 0 {
|
||
return v
|
||
}
|
||
return s.cfg.Default.UserBalance
|
||
}
|
||
|
||
// GetDefaultUserRPMLimit 获取新用户默认 RPM 限制(0 = 不限制)。未配置则返回 0。
|
||
func (s *SettingService) GetDefaultUserRPMLimit(ctx context.Context) int {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyDefaultUserRPMLimit)
|
||
if err != nil || value == "" {
|
||
return 0
|
||
}
|
||
if v, err := strconv.Atoi(value); err == nil && v >= 0 {
|
||
return v
|
||
}
|
||
return 0
|
||
}
|
||
|
||
// GetDefaultSubscriptions 获取新用户默认订阅配置列表。
|
||
func (s *SettingService) GetDefaultSubscriptions(ctx context.Context) []DefaultSubscriptionSetting {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyDefaultSubscriptions)
|
||
if err != nil {
|
||
return nil
|
||
}
|
||
return parseDefaultSubscriptions(value)
|
||
}
|
||
|
||
func (s *SettingService) GetAuthSourceDefaultSettings(ctx context.Context) (*AuthSourceDefaultSettings, error) {
|
||
keys := []string{
|
||
SettingKeyAuthSourceDefaultEmailBalance,
|
||
SettingKeyAuthSourceDefaultEmailConcurrency,
|
||
SettingKeyAuthSourceDefaultEmailSubscriptions,
|
||
SettingKeyAuthSourceDefaultEmailGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultEmailGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultLinuxDoBalance,
|
||
SettingKeyAuthSourceDefaultLinuxDoConcurrency,
|
||
SettingKeyAuthSourceDefaultLinuxDoSubscriptions,
|
||
SettingKeyAuthSourceDefaultLinuxDoGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultLinuxDoGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultOIDCBalance,
|
||
SettingKeyAuthSourceDefaultOIDCConcurrency,
|
||
SettingKeyAuthSourceDefaultOIDCSubscriptions,
|
||
SettingKeyAuthSourceDefaultOIDCGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultOIDCGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultWeChatBalance,
|
||
SettingKeyAuthSourceDefaultWeChatConcurrency,
|
||
SettingKeyAuthSourceDefaultWeChatSubscriptions,
|
||
SettingKeyAuthSourceDefaultWeChatGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultWeChatGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultGitHubBalance,
|
||
SettingKeyAuthSourceDefaultGitHubConcurrency,
|
||
SettingKeyAuthSourceDefaultGitHubSubscriptions,
|
||
SettingKeyAuthSourceDefaultGitHubGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultGitHubGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultGoogleBalance,
|
||
SettingKeyAuthSourceDefaultGoogleConcurrency,
|
||
SettingKeyAuthSourceDefaultGoogleSubscriptions,
|
||
SettingKeyAuthSourceDefaultGoogleGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultGoogleGrantOnFirstBind,
|
||
SettingKeyAuthSourceDefaultDingTalkBalance,
|
||
SettingKeyAuthSourceDefaultDingTalkConcurrency,
|
||
SettingKeyAuthSourceDefaultDingTalkSubscriptions,
|
||
SettingKeyAuthSourceDefaultDingTalkGrantOnSignup,
|
||
SettingKeyAuthSourceDefaultDingTalkGrantOnFirstBind,
|
||
SettingKeyAuthSourcePlatformQuotas("email"),
|
||
SettingKeyAuthSourcePlatformQuotas("linuxdo"),
|
||
SettingKeyAuthSourcePlatformQuotas("oidc"),
|
||
SettingKeyAuthSourcePlatformQuotas("wechat"),
|
||
SettingKeyAuthSourcePlatformQuotas("github"),
|
||
SettingKeyAuthSourcePlatformQuotas("google"),
|
||
SettingKeyAuthSourcePlatformQuotas("dingtalk"),
|
||
SettingKeyForceEmailOnThirdPartySignup,
|
||
}
|
||
|
||
settings, err := s.settingRepo.GetMultiple(ctx, keys)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("get auth source default settings: %w", err)
|
||
}
|
||
|
||
return &AuthSourceDefaultSettings{
|
||
Email: parseProviderDefaultGrantSettings(settings, emailAuthSourceDefaultKeys),
|
||
LinuxDo: parseProviderDefaultGrantSettings(settings, linuxDoAuthSourceDefaultKeys),
|
||
OIDC: parseProviderDefaultGrantSettings(settings, oidcAuthSourceDefaultKeys),
|
||
WeChat: parseProviderDefaultGrantSettings(settings, weChatAuthSourceDefaultKeys),
|
||
GitHub: parseProviderDefaultGrantSettings(settings, gitHubAuthSourceDefaultKeys),
|
||
Google: parseProviderDefaultGrantSettings(settings, googleAuthSourceDefaultKeys),
|
||
DingTalk: parseProviderDefaultGrantSettings(settings, dingTalkAuthSourceDefaultKeys),
|
||
ForceEmailOnThirdPartySignup: settings[SettingKeyForceEmailOnThirdPartySignup] == "true",
|
||
}, nil
|
||
}
|
||
|
||
func (s *SettingService) ResolveAuthSourceGrantSettings(ctx context.Context, signupSource string, firstBind bool) (ProviderDefaultGrantSettings, bool, error) {
|
||
result := ProviderDefaultGrantSettings{
|
||
Balance: s.GetDefaultBalance(ctx),
|
||
Concurrency: s.GetDefaultConcurrency(ctx),
|
||
Subscriptions: s.GetDefaultSubscriptions(ctx),
|
||
}
|
||
|
||
defaults, err := s.GetAuthSourceDefaultSettings(ctx)
|
||
if err != nil {
|
||
return result, false, err
|
||
}
|
||
|
||
providerDefaults, ok := authSourceSignupSettings(defaults, signupSource)
|
||
if !ok {
|
||
return result, false, nil
|
||
}
|
||
|
||
enabled := providerDefaults.GrantOnSignup
|
||
if firstBind {
|
||
enabled = providerDefaults.GrantOnFirstBind
|
||
}
|
||
if !enabled {
|
||
return result, false, nil
|
||
}
|
||
|
||
return mergeProviderDefaultGrantSettings(result, providerDefaults), true, nil
|
||
}
|
||
|
||
func (s *SettingService) UpdateAuthSourceDefaultSettings(ctx context.Context, settings *AuthSourceDefaultSettings) error {
|
||
updates, err := s.buildAuthSourceDefaultUpdates(ctx, settings)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
if len(updates) == 0 {
|
||
return nil
|
||
}
|
||
|
||
if err := s.settingRepo.SetMultiple(ctx, updates); err != nil {
|
||
return fmt.Errorf("update auth source default settings: %w", err)
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// IsTurnstileEnabled 检查是否启用 Turnstile 验证
|
||
func (s *SettingService) IsTurnstileEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyTurnstileEnabled)
|
||
if err != nil {
|
||
return false
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetTurnstileSecretKey 获取 Turnstile Secret Key
|
||
func (s *SettingService) GetTurnstileSecretKey(ctx context.Context) string {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyTurnstileSecretKey)
|
||
if err != nil {
|
||
return ""
|
||
}
|
||
return value
|
||
}
|
||
|
||
// IsIdentityPatchEnabled 检查是否启用身份补丁(Claude -> Gemini systemInstruction 注入)
|
||
func (s *SettingService) IsIdentityPatchEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyEnableIdentityPatch)
|
||
if err != nil {
|
||
// 默认开启,保持兼容
|
||
return true
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetIdentityPatchPrompt 获取自定义身份补丁提示词(为空表示使用内置默认模板)
|
||
func (s *SettingService) GetIdentityPatchPrompt(ctx context.Context) string {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyIdentityPatchPrompt)
|
||
if err != nil {
|
||
return ""
|
||
}
|
||
return value
|
||
}
|
||
|
||
// GenerateAdminAPIKey 生成新的管理员 API Key
|
||
func (s *SettingService) GenerateAdminAPIKey(ctx context.Context) (string, error) {
|
||
// 生成 32 字节随机数 = 64 位十六进制字符
|
||
bytes := make([]byte, 32)
|
||
if _, err := rand.Read(bytes); err != nil {
|
||
return "", fmt.Errorf("generate random bytes: %w", err)
|
||
}
|
||
|
||
key := AdminAPIKeyPrefix + hex.EncodeToString(bytes)
|
||
|
||
// 存储到 settings 表
|
||
if err := s.settingRepo.Set(ctx, SettingKeyAdminAPIKey, key); err != nil {
|
||
return "", fmt.Errorf("save admin api key: %w", err)
|
||
}
|
||
|
||
return key, nil
|
||
}
|
||
|
||
// GetAdminAPIKeyStatus 获取管理员 API Key 状态
|
||
// 返回脱敏的 key、是否存在、错误
|
||
func (s *SettingService) GetAdminAPIKeyStatus(ctx context.Context) (maskedKey string, exists bool, err error) {
|
||
key, err := s.settingRepo.GetValue(ctx, SettingKeyAdminAPIKey)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return "", false, nil
|
||
}
|
||
return "", false, err
|
||
}
|
||
if key == "" {
|
||
return "", false, nil
|
||
}
|
||
|
||
// 脱敏:显示前 10 位和后 4 位
|
||
if len(key) > 14 {
|
||
maskedKey = key[:10] + "..." + key[len(key)-4:]
|
||
} else {
|
||
maskedKey = key
|
||
}
|
||
|
||
return maskedKey, true, nil
|
||
}
|
||
|
||
// GetAdminAPIKey 获取完整的管理员 API Key(仅供内部验证使用)
|
||
// 如果未配置返回空字符串和 nil 错误,只有数据库错误时才返回 error
|
||
func (s *SettingService) GetAdminAPIKey(ctx context.Context) (string, error) {
|
||
key, err := s.settingRepo.GetValue(ctx, SettingKeyAdminAPIKey)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return "", nil // 未配置,返回空字符串
|
||
}
|
||
return "", err // 数据库错误
|
||
}
|
||
return key, nil
|
||
}
|
||
|
||
// DeleteAdminAPIKey 删除管理员 API Key
|
||
func (s *SettingService) DeleteAdminAPIKey(ctx context.Context) error {
|
||
return s.settingRepo.Delete(ctx, SettingKeyAdminAPIKey)
|
||
}
|
||
|
||
// IsModelFallbackEnabled 检查是否启用模型兜底机制
|
||
func (s *SettingService) IsModelFallbackEnabled(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyEnableModelFallback)
|
||
if err != nil {
|
||
return false // Default: disabled
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetFallbackModel 获取指定平台的兜底模型
|
||
func (s *SettingService) GetFallbackModel(ctx context.Context, platform string) string {
|
||
var key string
|
||
var defaultModel string
|
||
|
||
switch platform {
|
||
case PlatformAnthropic:
|
||
key = SettingKeyFallbackModelAnthropic
|
||
defaultModel = "claude-3-5-sonnet-20241022"
|
||
case PlatformOpenAI:
|
||
key = SettingKeyFallbackModelOpenAI
|
||
defaultModel = "gpt-4o"
|
||
case PlatformGemini:
|
||
key = SettingKeyFallbackModelGemini
|
||
defaultModel = "gemini-2.5-pro"
|
||
case PlatformAntigravity:
|
||
key = SettingKeyFallbackModelAntigravity
|
||
defaultModel = "gemini-2.5-pro"
|
||
default:
|
||
return ""
|
||
}
|
||
|
||
value, err := s.settingRepo.GetValue(ctx, key)
|
||
if err != nil || value == "" {
|
||
return defaultModel
|
||
}
|
||
return value
|
||
}
|
||
|
||
// GetOverloadCooldownSettings 获取529过载冷却配置
|
||
func (s *SettingService) GetOverloadCooldownSettings(ctx context.Context) (*OverloadCooldownSettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyOverloadCooldownSettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultOverloadCooldownSettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get overload cooldown settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultOverloadCooldownSettings(), nil
|
||
}
|
||
|
||
var settings OverloadCooldownSettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
return DefaultOverloadCooldownSettings(), nil
|
||
}
|
||
|
||
// 修正配置值范围
|
||
if settings.CooldownMinutes < 1 {
|
||
settings.CooldownMinutes = 1
|
||
}
|
||
if settings.CooldownMinutes > 120 {
|
||
settings.CooldownMinutes = 120
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// SetOverloadCooldownSettings 设置529过载冷却配置
|
||
func (s *SettingService) SetOverloadCooldownSettings(ctx context.Context, settings *OverloadCooldownSettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
// 禁用时修正为合法值即可,不拒绝请求
|
||
if settings.CooldownMinutes < 1 || settings.CooldownMinutes > 120 {
|
||
if settings.Enabled {
|
||
return fmt.Errorf("cooldown_minutes must be between 1-120")
|
||
}
|
||
settings.CooldownMinutes = 10 // 禁用状态下归一化为默认值
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal overload cooldown settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyOverloadCooldownSettings, string(data))
|
||
}
|
||
|
||
// GetRateLimit429CooldownSettings 获取429默认回避配置
|
||
func (s *SettingService) GetRateLimit429CooldownSettings(ctx context.Context) (*RateLimit429CooldownSettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyRateLimit429CooldownSettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultRateLimit429CooldownSettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get 429 cooldown settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultRateLimit429CooldownSettings(), nil
|
||
}
|
||
|
||
var settings RateLimit429CooldownSettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
return DefaultRateLimit429CooldownSettings(), nil
|
||
}
|
||
|
||
if settings.CooldownSeconds < 1 {
|
||
settings.CooldownSeconds = 1
|
||
}
|
||
if settings.CooldownSeconds > 7200 {
|
||
settings.CooldownSeconds = 7200
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// SetRateLimit429CooldownSettings 设置429默认回避配置
|
||
func (s *SettingService) SetRateLimit429CooldownSettings(ctx context.Context, settings *RateLimit429CooldownSettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
if settings.CooldownSeconds < 1 || settings.CooldownSeconds > 7200 {
|
||
if settings.Enabled {
|
||
return fmt.Errorf("cooldown_seconds must be between 1-7200")
|
||
}
|
||
settings.CooldownSeconds = 5
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal 429 cooldown settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyRateLimit429CooldownSettings, string(data))
|
||
}
|
||
|
||
// GetStreamTimeoutSettings 获取流超时处理配置
|
||
func (s *SettingService) GetStreamTimeoutSettings(ctx context.Context) (*StreamTimeoutSettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyStreamTimeoutSettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultStreamTimeoutSettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get stream timeout settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultStreamTimeoutSettings(), nil
|
||
}
|
||
|
||
var settings StreamTimeoutSettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
return DefaultStreamTimeoutSettings(), nil
|
||
}
|
||
|
||
// 验证并修正配置值
|
||
if settings.TempUnschedMinutes < 1 {
|
||
settings.TempUnschedMinutes = 1
|
||
}
|
||
if settings.TempUnschedMinutes > 60 {
|
||
settings.TempUnschedMinutes = 60
|
||
}
|
||
if settings.ThresholdCount < 1 {
|
||
settings.ThresholdCount = 1
|
||
}
|
||
if settings.ThresholdCount > 10 {
|
||
settings.ThresholdCount = 10
|
||
}
|
||
if settings.ThresholdWindowMinutes < 1 {
|
||
settings.ThresholdWindowMinutes = 1
|
||
}
|
||
if settings.ThresholdWindowMinutes > 60 {
|
||
settings.ThresholdWindowMinutes = 60
|
||
}
|
||
|
||
// 验证 action
|
||
switch settings.Action {
|
||
case StreamTimeoutActionTempUnsched, StreamTimeoutActionError, StreamTimeoutActionNone:
|
||
// valid
|
||
default:
|
||
settings.Action = StreamTimeoutActionTempUnsched
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// IsUngroupedKeySchedulingAllowed 查询是否允许未分组 Key 调度
|
||
func (s *SettingService) IsUngroupedKeySchedulingAllowed(ctx context.Context) bool {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyAllowUngroupedKeyScheduling)
|
||
if err != nil {
|
||
return false // fail-closed: 查询失败时默认不允许
|
||
}
|
||
return value == "true"
|
||
}
|
||
|
||
// GetRectifierSettings 获取请求整流器配置
|
||
func (s *SettingService) GetRectifierSettings(ctx context.Context) (*RectifierSettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyRectifierSettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultRectifierSettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get rectifier settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultRectifierSettings(), nil
|
||
}
|
||
|
||
var settings RectifierSettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
return DefaultRectifierSettings(), nil
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// SetRectifierSettings 设置请求整流器配置
|
||
func (s *SettingService) SetRectifierSettings(ctx context.Context, settings *RectifierSettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal rectifier settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyRectifierSettings, string(data))
|
||
}
|
||
|
||
// IsSignatureRectifierEnabled 判断签名整流是否启用(总开关 && 签名子开关)
|
||
func (s *SettingService) IsSignatureRectifierEnabled(ctx context.Context) bool {
|
||
settings, err := s.GetRectifierSettings(ctx)
|
||
if err != nil {
|
||
return true // fail-open: 查询失败时默认启用
|
||
}
|
||
return settings.Enabled && settings.ThinkingSignatureEnabled
|
||
}
|
||
|
||
// IsBudgetRectifierEnabled 判断 Budget 整流是否启用(总开关 && Budget 子开关)
|
||
func (s *SettingService) IsBudgetRectifierEnabled(ctx context.Context) bool {
|
||
settings, err := s.GetRectifierSettings(ctx)
|
||
if err != nil {
|
||
return true // fail-open: 查询失败时默认启用
|
||
}
|
||
return settings.Enabled && settings.ThinkingBudgetEnabled
|
||
}
|
||
|
||
// GetBetaPolicySettings 获取 Beta 策略配置
|
||
func (s *SettingService) GetBetaPolicySettings(ctx context.Context) (*BetaPolicySettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyBetaPolicySettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultBetaPolicySettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get beta policy settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultBetaPolicySettings(), nil
|
||
}
|
||
|
||
var settings BetaPolicySettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
return DefaultBetaPolicySettings(), nil
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// SetBetaPolicySettings 设置 Beta 策略配置
|
||
func (s *SettingService) SetBetaPolicySettings(ctx context.Context, settings *BetaPolicySettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
validActions := map[string]bool{
|
||
BetaPolicyActionPass: true, BetaPolicyActionFilter: true, BetaPolicyActionBlock: true,
|
||
}
|
||
validScopes := map[string]bool{
|
||
BetaPolicyScopeAll: true, BetaPolicyScopeOAuth: true, BetaPolicyScopeAPIKey: true, BetaPolicyScopeBedrock: true,
|
||
}
|
||
|
||
for i, rule := range settings.Rules {
|
||
if rule.BetaToken == "" {
|
||
return fmt.Errorf("rule[%d]: beta_token cannot be empty", i)
|
||
}
|
||
if !validActions[rule.Action] {
|
||
return fmt.Errorf("rule[%d]: invalid action %q", i, rule.Action)
|
||
}
|
||
if !validScopes[rule.Scope] {
|
||
return fmt.Errorf("rule[%d]: invalid scope %q", i, rule.Scope)
|
||
}
|
||
// Validate model_whitelist patterns
|
||
for j, pattern := range rule.ModelWhitelist {
|
||
trimmed := strings.TrimSpace(pattern)
|
||
if trimmed == "" {
|
||
return fmt.Errorf("rule[%d]: model_whitelist[%d] cannot be empty", i, j)
|
||
}
|
||
settings.Rules[i].ModelWhitelist[j] = trimmed
|
||
}
|
||
// Validate fallback_action
|
||
if rule.FallbackAction != "" && !validActions[rule.FallbackAction] {
|
||
return fmt.Errorf("rule[%d]: invalid fallback_action %q", i, rule.FallbackAction)
|
||
}
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal beta policy settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyBetaPolicySettings, string(data))
|
||
}
|
||
|
||
// GetOpenAIFastPolicySettings 获取 OpenAI fast 策略配置
|
||
func (s *SettingService) GetOpenAIFastPolicySettings(ctx context.Context) (*OpenAIFastPolicySettings, error) {
|
||
value, err := s.settingRepo.GetValue(ctx, SettingKeyOpenAIFastPolicySettings)
|
||
if err != nil {
|
||
if errors.Is(err, ErrSettingNotFound) {
|
||
return DefaultOpenAIFastPolicySettings(), nil
|
||
}
|
||
return nil, fmt.Errorf("get openai fast policy settings: %w", err)
|
||
}
|
||
if value == "" {
|
||
return DefaultOpenAIFastPolicySettings(), nil
|
||
}
|
||
|
||
var settings OpenAIFastPolicySettings
|
||
if err := json.Unmarshal([]byte(value), &settings); err != nil {
|
||
// JSON 损坏时静默 fallback 到默认配置会让策略意外失效(管理员配
|
||
// 置的 block/filter 规则被忽略)。记录 Warn 让运维能在出现异常
|
||
// 行为时定位到 settings 表里的脏数据。
|
||
slog.Warn("failed to unmarshal openai fast policy settings, falling back to defaults",
|
||
"error", err,
|
||
"key", SettingKeyOpenAIFastPolicySettings)
|
||
return DefaultOpenAIFastPolicySettings(), nil
|
||
}
|
||
|
||
return &settings, nil
|
||
}
|
||
|
||
// SetOpenAIFastPolicySettings 设置 OpenAI fast 策略配置
|
||
func (s *SettingService) SetOpenAIFastPolicySettings(ctx context.Context, settings *OpenAIFastPolicySettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
validActions := map[string]bool{
|
||
BetaPolicyActionPass: true, BetaPolicyActionFilter: true, BetaPolicyActionBlock: true,
|
||
OpenAIFastPolicyActionForcePriority: true,
|
||
}
|
||
validScopes := map[string]bool{
|
||
BetaPolicyScopeAll: true, BetaPolicyScopeOAuth: true, BetaPolicyScopeAPIKey: true, BetaPolicyScopeBedrock: true,
|
||
}
|
||
validTiers := map[string]bool{
|
||
OpenAIFastTierAny: true, OpenAIFastTierPriority: true, OpenAIFastTierFlex: true,
|
||
}
|
||
|
||
for i, rule := range settings.Rules {
|
||
tier := strings.ToLower(strings.TrimSpace(rule.ServiceTier))
|
||
if tier == "" {
|
||
tier = OpenAIFastTierAny
|
||
}
|
||
if !validTiers[tier] {
|
||
return fmt.Errorf("rule[%d]: invalid service_tier %q", i, rule.ServiceTier)
|
||
}
|
||
settings.Rules[i].ServiceTier = tier
|
||
if !validActions[rule.Action] {
|
||
return fmt.Errorf("rule[%d]: invalid action %q", i, rule.Action)
|
||
}
|
||
if !validScopes[rule.Scope] {
|
||
return fmt.Errorf("rule[%d]: invalid scope %q", i, rule.Scope)
|
||
}
|
||
for j, pattern := range rule.ModelWhitelist {
|
||
trimmed := strings.TrimSpace(pattern)
|
||
if trimmed == "" {
|
||
return fmt.Errorf("rule[%d]: model_whitelist[%d] cannot be empty", i, j)
|
||
}
|
||
settings.Rules[i].ModelWhitelist[j] = trimmed
|
||
}
|
||
if rule.FallbackAction != "" && !validActions[rule.FallbackAction] {
|
||
return fmt.Errorf("rule[%d]: invalid fallback_action %q", i, rule.FallbackAction)
|
||
}
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal openai fast policy settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyOpenAIFastPolicySettings, string(data))
|
||
}
|
||
|
||
// SetStreamTimeoutSettings 设置流超时处理配置
|
||
func (s *SettingService) SetStreamTimeoutSettings(ctx context.Context, settings *StreamTimeoutSettings) error {
|
||
if settings == nil {
|
||
return fmt.Errorf("settings cannot be nil")
|
||
}
|
||
|
||
// 验证配置值
|
||
if settings.TempUnschedMinutes < 1 || settings.TempUnschedMinutes > 60 {
|
||
return fmt.Errorf("temp_unsched_minutes must be between 1-60")
|
||
}
|
||
if settings.ThresholdCount < 1 || settings.ThresholdCount > 10 {
|
||
return fmt.Errorf("threshold_count must be between 1-10")
|
||
}
|
||
if settings.ThresholdWindowMinutes < 1 || settings.ThresholdWindowMinutes > 60 {
|
||
return fmt.Errorf("threshold_window_minutes must be between 1-60")
|
||
}
|
||
|
||
switch settings.Action {
|
||
case StreamTimeoutActionTempUnsched, StreamTimeoutActionError, StreamTimeoutActionNone:
|
||
// valid
|
||
default:
|
||
return fmt.Errorf("invalid action: %s", settings.Action)
|
||
}
|
||
|
||
data, err := json.Marshal(settings)
|
||
if err != nil {
|
||
return fmt.Errorf("marshal stream timeout settings: %w", err)
|
||
}
|
||
|
||
return s.settingRepo.Set(ctx, SettingKeyStreamTimeoutSettings, string(data))
|
||
}
|
||
|
||
// GetDefaultPlatformQuotas 读取系统全局 platform quota JSON key,返回全部允许平台 x 3 window 的设置。
|
||
// 永远返回包含全部允许 platform key 的 map(值可能为零值/nil 字段,表示"上层未配置 = 不限制")。
|
||
//
|
||
// 使用单个 JSON key(default_platform_quotas),一次 DB roundtrip,消除旧 12-KV 格式的 N+1 问题。
|
||
// 容错语义:取值失败或 unmarshal 失败 → 返回补齐全部允许平台 key 的空 map(fail-open,注册不被阻断)。
|
||
func (s *SettingService) GetDefaultPlatformQuotas(ctx context.Context) (map[string]*DefaultPlatformQuotaSetting, error) {
|
||
out := make(map[string]*DefaultPlatformQuotaSetting, len(AllowedQuotaPlatforms))
|
||
for _, platform := range AllowedQuotaPlatforms {
|
||
out[platform] = &DefaultPlatformQuotaSetting{}
|
||
}
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyDefaultPlatformQuotas)
|
||
if err != nil || raw == "" {
|
||
return out, nil // 无配置 = 全部不限制
|
||
}
|
||
parsed := map[string]*DefaultPlatformQuotaSetting{}
|
||
if err := json.Unmarshal([]byte(raw), &parsed); err != nil {
|
||
slog.Warn("[Setting] unmarshal default_platform_quotas failed (fail-open)", "error", err)
|
||
return out, nil
|
||
}
|
||
for _, platform := range AllowedQuotaPlatforms {
|
||
if v := parsed[platform]; v != nil {
|
||
out[platform] = v
|
||
}
|
||
}
|
||
return out, nil // 补齐全部允许 platform key,保持与旧实现一致的下游契约
|
||
}
|
||
|
||
// GetAuthSourcePlatformQuotas 读取指定 auth source 的 platform quota 覆盖(仅返回有配置的平台,override 语义)。
|
||
func (s *SettingService) GetAuthSourcePlatformQuotas(ctx context.Context, source string) map[string]*DefaultPlatformQuotaSetting {
|
||
out := map[string]*DefaultPlatformQuotaSetting{}
|
||
raw, err := s.settingRepo.GetValue(ctx, SettingKeyAuthSourcePlatformQuotas(source))
|
||
if err != nil || raw == "" {
|
||
return out // 无 override
|
||
}
|
||
if err := json.Unmarshal([]byte(raw), &out); err != nil {
|
||
slog.Warn("[Setting] unmarshal auth source platform quotas failed (fail-open)", "source", source, "error", err)
|
||
return map[string]*DefaultPlatformQuotaSetting{}
|
||
}
|
||
return out // 仅含已配置平台,保持 override 语义
|
||
}
|
||
|
||
// mergePlatformQuotaDefaults 按字段级 patch:src 中非 nil 字段覆盖 dst。
|
||
// 区分 nil("未配置",保留 dst)vs &0.0("显式禁用",覆盖 dst 为 0)
|
||
func mergePlatformQuotaDefaults(dst, src *DefaultPlatformQuotaSetting) {
|
||
if src == nil || dst == nil {
|
||
return
|
||
}
|
||
if src.DailyLimitUSD != nil {
|
||
dst.DailyLimitUSD = src.DailyLimitUSD
|
||
}
|
||
if src.WeeklyLimitUSD != nil {
|
||
dst.WeeklyLimitUSD = src.WeeklyLimitUSD
|
||
}
|
||
if src.MonthlyLimitUSD != nil {
|
||
dst.MonthlyLimitUSD = src.MonthlyLimitUSD
|
||
}
|
||
}
|