62 Коміти

Автор SHA1 Повідомлення Дата
  fengsilin caf457ce9e feat: 错误时记录上游响应体,支持流式 3 місяці тому
  fengsilin 575f84064e feat: codex API key 模式 + 前端凭证回显修复 3 місяці тому
  fengsilin ca91922904 docs: add log chat/upstream id design 3 місяці тому
  fengsilin 332d5a62f2 feat: support codex api key credentials 3 місяці тому
  fengsilin 1c0898257c feat: add redemption remarks 3 місяці тому
  fengsilin f18e51e3a9 feat: add redemption remark support 3 місяці тому
  fengsilin fae24fdf37 test: stabilize channel affinity usage cache tests 3 місяці тому
  fengsilin fc0257f6ca chore: ignore local worktrees 3 місяці тому
  fengsilin d8272e7707 feat(channel): 添加渠道"对外名称"(public_name)字段 3 місяці тому
  fengsilin 73d10b5799 feat: 错误日志记录上游 request-id 和响应体,Playground 渠道路由改为 header 传递 3 місяці тому
  fengsilin 5b8303ae81 fix: 登录后跳转来源页,充值金额校验优化 3 місяці тому
  fengsilin 8bb2883d7e feat(home): 首页定价卡片添加"去体验"按钮,跳转 Playground 3 місяці тому
  fengsilin a0ac37b2a9 fix: 替换 println 为 SysLog,修复 logger nil context 崩溃 3 місяці тому
  fengsilin 74cc8c0d56 feat(playground): 添加渠道选择功能,支持指定渠道体验模型 3 місяці тому
  fengsilin a34d44a824 fix(user): 修复从节点同步用户余额显示为 0 的问题 3 місяці тому
  fengsilin 967309fb56 feat: 添加邮箱后缀注册额度规则功能 3 місяці тому
  fengsilin 83f767b40b Merge branch 'worktree-feat-captcha' 3 місяці тому
  fengsilin d24d68a916 fix(i18n): 修复令牌和充值页面货币显示不一致问题 3 місяці тому
  fengsilin fe88bc7758 Merge branch 'worktree-feat-captcha' 3 місяці тому
  fengsilin 1912dd3b72 feat: 添加图片验证码功能,防止脚本批量注册 3 місяці тому
  fengsilin 20163dbad5 fix(legal): 法律文档页面切换语言时强制重新加载内容 3 місяці тому
  fengsilin 56e8d14485 fix: 语言切换后 TextArea 内容丢失 3 місяці тому
  fengsilin 65bd1dabb8 fix: 法律文档编辑器切换语言时 TextArea 内容未刷新 3 місяці тому
  fengsilin 934b4b10a7 feat: 法律文档双语配置(中/英)+ 精简语言支持至中英双语 3 місяці тому
  fengsilin 676902ed58 Merge branch 'worktree-feat-terms-usage-policy' 3 місяці тому
  fengsilin 2757f77011 feat(i18n): 补全服务条款和使用政策的多语言翻译 3 місяці тому
  fengsilin 65c6446a30 Merge branch 'worktree-feat-terms-usage-policy' 3 місяці тому
  fengsilin 1cc75aa9a2 feat: 服务条款和使用政策页面,支持后台 Markdown 配置 3 місяці тому
  fengsilin ed195cada5 fix(frontend): sort_order 默认值显示为"未设置",编辑弹窗增加提示 3 місяці тому
  fengsilin d5157a779c fix(pricing): 统一 sort_order 排序逻辑,有 meta 但未设置的不优先 3 місяці тому
  fengsilin 54e53108ec fix(sort): sort_order 默认值改为 999999,简化排序逻辑 3 місяці тому
  fengsilin 3b99c83e32 fix(sort): sort_order=0 的记录排到最后,非0值按升序排列 3 місяці тому
  fengsilin 518f1ab87f refactor(frontend): 移除拖拽排序,改为 sort_order 数值输入 3 місяці тому
  fengsilin 65880fc69b feat(frontend): Vendor Tab 和 Model 表格支持拖拽排序,移除定价页字母排序 3 місяці тому
  fengsilin 5e0f38d26c chore(frontend): 安装 @dnd-kit 拖拽排序库 3 місяці тому
  fengsilin e507897f21 feat(pricing): updatePricing 按 sort_order 排序 vendors 和 models 3 місяці тому
  fengsilin 749064932c feat: 新增 PUT /api/vendors/reorder 和 /api/models/reorder 批量排序接口 3 місяці тому
  fengsilin f8859b7c35 feat: Vendor 和 Model 添加 sort_order 字段,查询按 sort_order ASC 排序 3 місяці тому
  fengsilin 3a29f3772f feat(channel): 添加模型默认通道功能,支持优先路由和卡片标识 3 місяці тому
  fengsilin 93e331b624 feat: 支持 LOGO_FILE_PATH 环境变量指定本地 Logo,优化定价与语言设置 3 місяці тому
  fengsilin 94f24b29f7 style: 首页文案"全球"改为"顶级" 3 місяці тому
  fengsilin f0f0215b20 fix(pricing): 修复缓存倍率写入默认值1和缓存创建token丢失 3 місяці тому
  fengsilin 387f6c1ae0 feat(pricing): 定价数据源切换到渠道表,新增缓存价格展示 3 місяці тому
  fengsilin 00058cd671 refactor(channel-pricing): 提取辅助方法消除重复代码,修复缓存一致性 3 місяці тому
  fengsilin 42082311d3 fix: 修复默认语言不生效的问题 3 місяці тому
  fengsilin 9eb58c684e feat: 后台设置默认语言 3 місяці тому
  fengsilin a2b27ea2a6 docs: 后台设置默认语言实施计划 3 місяці тому
  fengsilin 436fcb405b docs: 后台设置默认语言功能设计文档 3 місяці тому
  fengsilin 8d22c4e80e style(home): 临时隐藏首页工具链/核心价值/工作流/生态伙伴 section 3 місяці тому
  fengsilin 355757404b merge: feat/channel-pricing-extended → master 3 місяці тому
  fengsilin 235f7c6e5f refactor(channel-pricing): 提取 ParseTagIds 辅助函数 + 移除调试 console.log 3 місяці тому
  fengsilin f3577590bc test(channel-pricing): 添加 API 端到端集成测试脚本 3 місяці тому
  fengsilin 129e17ac42 feat(channel-pricing): 前端支持缓存/图片/音频倍率编辑和展示 3 місяці тому
  fengsilin c84c26b5a2 feat(channel-pricing): GetChannelPricingByModelWithChannelInfo 支持 CASE WHEN 回退 + 扩展字段 3 місяці тому
  fengsilin d5d4714908 test(channel-pricing): 缓存写穿 + 字段默认值单元测试 3 місяці тому
  fengsilin 646f37dd4b feat(channel-pricing): Controller 扩展 API 支持新字段 + 输入校验 + 操作日志 3 місяці тому
  fengsilin 3a24fd598f feat(channel-pricing): ModelPriceHelper + UpdatePriceDataForChannelPricing 适配新签名,支持扩展比率覆盖 3 місяці тому
  fengsilin e5f029f91e feat(channel-pricing): InitDB 启动时全量加载渠道定价缓存 3 місяці тому
  fengsilin 2d4c73d3aa refactor(channel-pricing): 结构体新增扩展字段 + 缓存重写为全量加载写穿 3 місяці тому
  fengsilin 156618fdad merge: feat/alipay-payment → master 3 місяці тому
  fengsilin b514a0798b refactor(payment): 合并微信/支付宝支付重复代码,删除冗余文件 3 місяці тому
  fengsilin e8ade3abae feat(payment): 集成支付宝当面付扫码支付 + 补全前端 i18n 硬编码中文 3 місяці тому
100 змінених файлів з 4624 додано та 510 видалено
  1. +21
    -1
      .dockerignore
  2. +1
    -0
      .gitattributes
  3. +1
    -0
      .gitignore
  4. +30
    -0
      common/captcha.go
  5. +58
    -0
      common/codex_credential.go
  6. +10
    -0
      common/constants.go
  7. +18
    -0
      common/str.go
  8. +28
    -16
      controller/channel.go
  9. +120
    -55
      controller/channel_pricing.go
  10. +111
    -0
      controller/codex_channel_test.go
  11. +8
    -10
      controller/codex_usage.go
  12. +128
    -0
      controller/email_quota_rule.go
  13. +264
    -0
      controller/email_quota_rule_test.go
  14. +74
    -5
      controller/misc.go
  15. +62
    -0
      controller/playground_channels.go
  16. +2
    -0
      controller/redemption.go
  17. +159
    -0
      controller/redemption_test.go
  18. +6
    -0
      controller/relay.go
  19. +55
    -0
      controller/reorder.go
  20. +47
    -4
      controller/topup.go
  21. +280
    -0
      controller/topup_alipay.go
  22. +2
    -24
      controller/topup_wechat.go
  23. +5
    -1
      controller/user.go
  24. +1
    -0
      docs/DATABASE_SCHEMA.md
  25. +307
    -0
      docs/superpowers/plans/2026-04-17-default-language-setting.md
  26. +114
    -0
      docs/superpowers/specs/2026-04-17-default-language-setting-design.md
  27. +236
    -0
      docs/superpowers/specs/2026-04-30-log-chatid-upstreamid-design.md
  28. +1
    -0
      dto/openai_response.go
  29. +2
    -0
      go.mod
  30. +57
    -0
      go.sum
  31. +5
    -3
      logger/logger.go
  32. +10
    -0
      main.go
  33. +62
    -14
      middleware/distributor.go
  34. +53
    -0
      model/ability.go
  35. +9
    -1
      model/channel.go
  36. +226
    -79
      model/channel_pricing.go
  37. +242
    -0
      model/channel_pricing_test.go
  38. +133
    -0
      model/email_quota_rule.go
  39. +246
    -0
      model/email_quota_rule_test.go
  40. +26
    -3
      model/main.go
  41. +29
    -3
      model/model_meta.go
  42. +33
    -0
      model/option.go
  43. +134
    -9
      model/pricing.go
  44. +122
    -0
      model/pricing_test.go
  45. +4
    -3
      model/redemption.go
  46. +101
    -14
      model/redemption_test.go
  47. +32
    -1
      model/topup.go
  48. +12
    -6
      model/topup_wechat.go
  49. +40
    -6
      model/user.go
  50. +18
    -2
      model/vendor_meta.go
  51. +2
    -2
      relay/channel/api_request.go
  52. +4
    -4
      relay/channel/claude/relay-claude.go
  53. +27
    -11
      relay/channel/codex/adaptor.go
  54. +5
    -9
      relay/channel/openai/chat_via_responses.go
  55. +11
    -4
      relay/channel/openai/relay_responses.go
  56. +29
    -0
      relay/channel/openai/responses_error.go
  57. +64
    -0
      relay/channel/openai/upstream_body_test.go
  58. +1
    -1
      relay/claude_handler.go
  59. +17
    -4
      relay/compatible_handler.go
  60. +60
    -41
      relay/helper/price.go
  61. +22
    -0
      router/api-router.go
  62. +8
    -0
      router/web-router.go
  63. +6
    -7
      service/channel_affinity_usage_cache_test.go
  64. +2
    -24
      service/codex_credential_refresh.go
  65. +1
    -1
      service/codex_credential_refresh_task.go
  66. +21
    -0
      service/error.go
  67. +32
    -0
      service/truncate_body_test.go
  68. +18
    -0
      setting/payment_alipay.go
  69. +18
    -0
      setting/ratio_setting/model_ratio.go
  70. +16
    -4
      setting/system_setting/legal.go
  71. +122
    -0
      test-scripts/test_channel_pricing_extended.sh
  72. +10
    -8
      types/error.go
  73. +26
    -0
      types/price_data.go
  74. +1
    -1
      web/index.html
  75. +18
    -0
      web/src/App.jsx
  76. +21
    -11
      web/src/components/auth/LoginForm.jsx
  77. +1
    -1
      web/src/components/auth/PasswordResetConfirm.jsx
  78. +1
    -1
      web/src/components/auth/PasswordResetForm.jsx
  79. +68
    -9
      web/src/components/auth/RegisterForm.jsx
  80. +3
    -2
      web/src/components/common/markdown/MarkdownRenderer.jsx
  81. +1
    -1
      web/src/components/common/ui/JSONEditor.jsx
  82. +2
    -0
      web/src/components/layout/Footer.jsx
  83. +5
    -4
      web/src/components/layout/PageLayout.jsx
  84. +2
    -0
      web/src/components/layout/headerbar/HeaderLogo.jsx
  85. +0
    -30
      web/src/components/layout/headerbar/LanguageSelector.jsx
  86. +1
    -1
      web/src/components/playground/CodeViewer.jsx
  87. +1
    -1
      web/src/components/playground/DebugPanel.jsx
  88. +2
    -2
      web/src/components/playground/MessageContent.jsx
  89. +2
    -1
      web/src/components/playground/OptimizedComponents.js
  90. +33
    -0
      web/src/components/playground/SettingsPanel.jsx
  91. +2
    -2
      web/src/components/playground/ThinkingContent.jsx
  92. +5
    -4
      web/src/components/playground/configStorage.js
  93. +1
    -1
      web/src/components/settings/ModelDeploymentSetting.jsx
  94. +1
    -1
      web/src/components/settings/ModelSetting.jsx
  95. +137
    -53
      web/src/components/settings/OtherSetting.jsx
  96. +6
    -0
      web/src/components/settings/PaymentSetting.jsx
  97. +1
    -1
      web/src/components/settings/RateLimitSetting.jsx
  98. +1
    -1
      web/src/components/settings/RatioSetting.jsx
  99. +40
    -1
      web/src/components/settings/SystemSetting.jsx
  100. +1
    -1
      web/src/components/settings/personal/cards/NotificationSettings.jsx

+ 21
- 1
.dockerignore Переглянути файл

@@ -7,4 +7,24 @@ Makefile
docs docs
.eslintcache .eslintcache
.gocache .gocache
/web/node_modules
/web/node_modules
.claude
.plans
current-page*
models-page*
pricing-page.png
login-page
scripts
relay/helper/price_test.go
.superpowers
*.png
*.bak
.worktrees
**/node_modules
**/.gocache
**/.gocache-temp
logs
*.db
*.db-journal
*.zip
web/dist

+ 1
- 0
.gitattributes Переглянути файл

@@ -36,3 +36,4 @@
# ============================================ # ============================================
# Mark web frontend as vendored so GitHub recognizes this as a Go project # Mark web frontend as vendored so GitHub recognizes this as a Go project
electron/** linguist-vendored electron/** linguist-vendored
.dockerignore text eol=lf

+ 1
- 0
.gitignore Переглянути файл

@@ -23,6 +23,7 @@ plans
docs/plans/ docs/plans/
CLAUDE.md CLAUDE.md
.claude .claude
.worktrees/
logs/ logs/
docs/superpowers docs/superpowers




+ 30
- 0
common/captcha.go Переглянути файл

@@ -0,0 +1,30 @@
package common

import (
"github.com/mojocn/base64Captcha"
)

var captchaStore = base64Captcha.DefaultMemStore

var captchaDriver = base64Captcha.NewDriverString(
40, // height
120, // width
0, // noise count (auto)
base64Captcha.OptionShowSlimeLine, // show slime lines
5, // code length
base64Captcha.TxtSimpleCharaters, // digits+letters excluding confusing chars
nil, // bg color (auto)
nil, // font storage (auto)
[]string{"wqy-microhei.ttc"}, // font files
)

var captchaInstance = base64Captcha.NewCaptcha(captchaDriver, captchaStore)

func GenerateCaptcha() (string, string, error) {
id, b64s, _, err := captchaInstance.Generate()
return id, b64s, err
}

func VerifyCaptcha(id, code string) bool {
return captchaStore.Verify(id, code, true)
}

+ 58
- 0
common/codex_credential.go Переглянути файл

@@ -0,0 +1,58 @@
package common

import (
"errors"
"strings"
)

type CodexCredentialMode string

const (
CodexCredentialModeAPIKey CodexCredentialMode = "api_key"
CodexCredentialModeOAuth CodexCredentialMode = "oauth"
)

var (
ErrCodexOAuthCredentialRequired = errors.New("codex channel: oauth credential required")
ErrCodexOAuthCredentialInvalidJSON = errors.New("codex channel: invalid oauth key json")
ErrCodexOAuthAccessTokenRequired = errors.New("codex channel: access_token is required")
ErrCodexOAuthAccountIDRequired = errors.New("codex channel: account_id is required")
)

type CodexOAuthCredential struct {
IDToken string `json:"id_token,omitempty"`
AccessToken string `json:"access_token,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
AccountID string `json:"account_id,omitempty"`
LastRefresh string `json:"last_refresh,omitempty"`
Email string `json:"email,omitempty"`
Type string `json:"type,omitempty"`
Expired string `json:"expired,omitempty"`
}

func ParseCodexOAuthCredential(raw string) (*CodexOAuthCredential, error) {
trimmed := strings.TrimSpace(raw)
if trimmed == "" || !strings.HasPrefix(trimmed, "{") {
return nil, ErrCodexOAuthCredentialRequired
}

var credential CodexOAuthCredential
if err := Unmarshal([]byte(trimmed), &credential); err != nil {
return nil, ErrCodexOAuthCredentialInvalidJSON
}
if strings.TrimSpace(credential.AccessToken) == "" {
return nil, ErrCodexOAuthAccessTokenRequired
}
if strings.TrimSpace(credential.AccountID) == "" {
return nil, ErrCodexOAuthAccountIDRequired
}

return &credential, nil
}

func DetectCodexCredentialMode(raw string) CodexCredentialMode {
if _, err := ParseCodexOAuthCredential(raw); err == nil {
return CodexCredentialModeOAuth
}
return CodexCredentialModeAPIKey
}

+ 10
- 0
common/constants.go Переглянути файл

@@ -15,7 +15,16 @@ var Version = "v0.0.0" // this hard coding will be replaced automatic
var SystemName = "New API" var SystemName = "New API"
var Footer = "" var Footer = ""
var Logo = "" var Logo = ""
var LogoFilePath = "" // LOGO_FILE_PATH 环境变量指定的本地 Logo 文件路径

func GetEffectiveLogo() string {
if LogoFilePath != "" {
return "/logo.png"
}
return Logo
}
var TopUpLink = "" var TopUpLink = ""
var DefaultLanguage = "" // admin-configured default language; empty = follow browser detection


// var ChatLink = "" // var ChatLink = ""
// var ChatLink2 = "" // var ChatLink2 = ""
@@ -49,6 +58,7 @@ var LinuxDOOAuthEnabled = false
var WeChatAuthEnabled = false var WeChatAuthEnabled = false
var TelegramOAuthEnabled = false var TelegramOAuthEnabled = false
var TurnstileCheckEnabled = false var TurnstileCheckEnabled = false
var CaptchaEnabled = false
var RegisterEnabled = true var RegisterEnabled = true


var EmailDomainRestrictionEnabled = false // 是否启用邮箱域名限制 var EmailDomainRestrictionEnabled = false // 是否启用邮箱域名限制


+ 18
- 0
common/str.go Переглянути файл

@@ -87,6 +87,24 @@ func StringsContains(strs []string, str string) bool {
return false return false
} }


// StringsSubtract returns elements from source that are not in exclude.
func StringsSubtract(source, exclude []string) []string {
if len(exclude) == 0 {
return source
}
excludeSet := make(map[string]struct{}, len(exclude))
for _, s := range exclude {
excludeSet[s] = struct{}{}
}
result := make([]string, 0, len(source))
for _, s := range source {
if _, ok := excludeSet[s]; !ok {
result = append(result, s)
}
}
return result
}

// StringToByteSlice []byte only read, panic on append // StringToByteSlice []byte only read, panic on append
func StringToByteSlice(s string) []byte { func StringToByteSlice(s string) []byte {
tmp1 := (*[2]uintptr)(unsafe.Pointer(&s)) tmp1 := (*[2]uintptr)(unsafe.Pointer(&s))


+ 28
- 16
controller/channel.go Переглянути файл

@@ -3,6 +3,7 @@ package controller
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strconv" "strconv"
@@ -584,6 +585,10 @@ func validateChannel(channel *model.Channel, isAdd bool) error {
return fmt.Errorf("channel cannot be empty") return fmt.Errorf("channel cannot be empty")
} }


if strings.TrimSpace(channel.PublicName) == "" {
return fmt.Errorf("public name cannot be empty")
}

// 检查模型名称长度是否超过 255 // 检查模型名称长度是否超过 255
for _, m := range channel.GetModels() { for _, m := range channel.GetModels() {
if len(m) > 255 { if len(m) > 255 {
@@ -611,20 +616,22 @@ func validateChannel(channel *model.Channel, isAdd bool) error {
// Codex OAuth key validation (optional, only when JSON object is provided) // Codex OAuth key validation (optional, only when JSON object is provided)
if channel.Type == constant.ChannelTypeCodex { if channel.Type == constant.ChannelTypeCodex {
trimmedKey := strings.TrimSpace(channel.Key) trimmedKey := strings.TrimSpace(channel.Key)
if isAdd || trimmedKey != "" {
if !strings.HasPrefix(trimmedKey, "{") {
return fmt.Errorf("Codex key must be a valid JSON object")
}
var keyMap map[string]any
if err := common.Unmarshal([]byte(trimmedKey), &keyMap); err != nil {
if isAdd && trimmedKey == "" {
return fmt.Errorf("Codex key cannot be empty")
}
if strings.HasPrefix(trimmedKey, "{") {
if _, err := common.ParseCodexOAuthCredential(trimmedKey); err != nil {
if errors.Is(err, common.ErrCodexOAuthCredentialInvalidJSON) {
return fmt.Errorf("Codex key must be a valid JSON object")
}
if errors.Is(err, common.ErrCodexOAuthAccessTokenRequired) {
return fmt.Errorf("Codex key JSON must include access_token")
}
if errors.Is(err, common.ErrCodexOAuthAccountIDRequired) {
return fmt.Errorf("Codex key JSON must include account_id")
}
return fmt.Errorf("Codex key must be a valid JSON object") return fmt.Errorf("Codex key must be a valid JSON object")
} }
if v, ok := keyMap["access_token"]; !ok || v == nil || strings.TrimSpace(fmt.Sprintf("%v", v)) == "" {
return fmt.Errorf("Codex key JSON must include access_token")
}
if v, ok := keyMap["account_id"]; !ok || v == nil || strings.TrimSpace(fmt.Sprintf("%v", v)) == "" {
return fmt.Errorf("Codex key JSON must include account_id")
}
} }
} }


@@ -643,6 +650,10 @@ func RefreshCodexChannelCredential(c *gin.Context) {


oauthKey, ch, err := service.RefreshCodexChannelCredential(ctx, channelId, service.CodexCredentialRefreshOptions{ResetCaches: true}) oauthKey, ch, err := service.RefreshCodexChannelCredential(ctx, channelId, service.CodexCredentialRefreshOptions{ResetCaches: true})
if err != nil { if err != nil {
if errors.Is(err, common.ErrCodexOAuthCredentialRequired) {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "当前凭证方式不支持刷新凭证"})
return
}
common.SysError("failed to refresh codex channel credential: " + err.Error()) common.SysError("failed to refresh codex channel credential: " + err.Error())
c.JSON(http.StatusOK, gin.H{"success": false, "message": "刷新凭证失败,请稍后重试"}) c.JSON(http.StatusOK, gin.H{"success": false, "message": "刷新凭证失败,请稍后重试"})
return return
@@ -2108,10 +2119,11 @@ func GetUserChannelsForBinding(c *gin.Context) {
result := make([]gin.H, 0, len(channels)) result := make([]gin.H, 0, len(channels))
for _, ch := range channels { for _, ch := range channels {
result = append(result, gin.H{ result = append(result, gin.H{
"id": ch.Id,
"name": ch.Name,
"type": ch.Type,
"remark": ch.Remark,
"id": ch.Id,
"name": ch.Name,
"public_name": ch.PublicName,
"type": ch.Type,
"remark": ch.Remark,
}) })
} }




+ 120
- 55
controller/channel_pricing.go Переглянути файл

@@ -1,6 +1,7 @@
package controller package controller


import ( import (
"fmt"
"strconv" "strconv"
"strings" "strings"


@@ -50,14 +51,25 @@ func GetChannelPricingByModel(c *gin.Context) {


// CreateChannelPricingRequest 创建渠道定价请求 // CreateChannelPricingRequest 创建渠道定价请求
type CreateChannelPricingRequest struct { type CreateChannelPricingRequest struct {
Id int `json:"id"`
ModelName string `json:"model_name" binding:"required"`
ChannelId int `json:"channel_id" binding:"required"`
QuotaType int `json:"quota_type"`
ModelRatio float64 `json:"model_ratio"`
CompletionRatio float64 `json:"completion_ratio"`
ModelPrice float64 `json:"model_price"`
TagIds string `json:"tag_ids"`
Id int `json:"id"`
ModelName string `json:"model_name" binding:"required"`
ChannelId int `json:"channel_id" binding:"required"`
QuotaType int `json:"quota_type"`
ModelRatio float64 `json:"model_ratio"`
CompletionRatio float64 `json:"completion_ratio"`
ModelPrice float64 `json:"model_price"`
TagIds string `json:"tag_ids"`
CacheRatio float64 `json:"cache_ratio"`
CacheCreationRatio float64 `json:"cache_creation_ratio"`
ImageRatio float64 `json:"image_ratio"`
AudioRatio float64 `json:"audio_ratio"`
AudioCompletionRatio float64 `json:"audio_completion_ratio"`
}

// applyRequest 将请求字段应用到 ChannelPricing
func applyRequestFields(cp *model.ChannelPricing, req *CreateChannelPricingRequest) {
cp.ApplyFields(req.QuotaType, req.ModelRatio, req.CompletionRatio, req.ModelPrice, req.TagIds,
req.CacheRatio, req.CacheCreationRatio, req.ImageRatio, req.AudioRatio, req.AudioCompletionRatio)
} }


// CreateChannelPricing 创建或更新渠道定价 // CreateChannelPricing 创建或更新渠道定价
@@ -68,38 +80,38 @@ func CreateChannelPricing(c *gin.Context) {
return return
} }


if req.CacheRatio < 0 || req.CacheCreationRatio < 0 || req.ImageRatio < 0 || req.AudioRatio < 0 || req.AudioCompletionRatio < 0 {
common.ApiErrorMsg(c, "ratio values must be >= 0")
return
}

// 检查是否已存在 // 检查是否已存在
existing, _ := model.GetChannelPricing(req.ModelName, req.ChannelId) existing, _ := model.GetChannelPricing(req.ModelName, req.ChannelId)
if existing != nil { if existing != nil {
// 更新
existing.QuotaType = req.QuotaType
existing.ModelRatio = req.ModelRatio
existing.CompletionRatio = req.CompletionRatio
existing.ModelPrice = req.ModelPrice
existing.TagIds = req.TagIds
applyRequestFields(existing, &req)
if err := existing.Update(); err != nil { if err := existing.Update(); err != nil {
common.ApiError(c, err) common.ApiError(c, err)
return return
} }
common.SysLog(fmt.Sprintf("[ChannelPricing] updated: id=%d model=%s channel=%d", existing.Id, existing.ModelName, existing.ChannelId))
common.ApiSuccess(c, existing) common.ApiSuccess(c, existing)
return return
} }


// 创建 // 创建
cp := &model.ChannelPricing{ cp := &model.ChannelPricing{
ModelName: req.ModelName,
ChannelId: req.ChannelId,
QuotaType: req.QuotaType,
ModelRatio: req.ModelRatio,
CompletionRatio: req.CompletionRatio,
ModelPrice: req.ModelPrice,
TagIds: req.TagIds,
ModelName: req.ModelName,
ChannelId: req.ChannelId,
} }
applyRequestFields(cp, &req)
if err := cp.Insert(); err != nil { if err := cp.Insert(); err != nil {
common.ApiError(c, err) common.ApiError(c, err)
return return
} }


common.SysLog(fmt.Sprintf("[ChannelPricing] created: model=%s channel=%d quotaType=%d modelRatio=%.4f completionRatio=%.4f modelPrice=%.4f cacheRatio=%.4f cacheCreationRatio=%.4f imageRatio=%.4f audioRatio=%.4f audioCompletionRatio=%.4f",
req.ModelName, req.ChannelId, req.QuotaType, req.ModelRatio, req.CompletionRatio, req.ModelPrice,
req.CacheRatio, req.CacheCreationRatio, req.ImageRatio, req.AudioRatio, req.AudioCompletionRatio))
common.ApiSuccess(c, cp) common.ApiSuccess(c, cp)
} }


@@ -118,26 +130,20 @@ func BatchCreateChannelPricing(c *gin.Context) {


pricings := make([]*model.ChannelPricing, 0, len(req.Items)) pricings := make([]*model.ChannelPricing, 0, len(req.Items))
for _, item := range req.Items { for _, item := range req.Items {
pricings = append(pricings, &model.ChannelPricing{
ModelName: item.ModelName,
ChannelId: item.ChannelId,
QuotaType: item.QuotaType,
ModelRatio: item.ModelRatio,
CompletionRatio: item.CompletionRatio,
ModelPrice: item.ModelPrice,
TagIds: item.TagIds,
})
cp := &model.ChannelPricing{
ModelName: item.ModelName,
ChannelId: item.ChannelId,
}
applyRequestFields(cp, item)
pricings = append(pricings, cp)
} }


// 使用事务逐个处理(GORM 的批量 upsert 在不同数据库表现不一致) // 使用事务逐个处理(GORM 的批量 upsert 在不同数据库表现不一致)
for _, cp := range pricings { for _, cp := range pricings {
existing, _ := model.GetChannelPricing(cp.ModelName, cp.ChannelId) existing, _ := model.GetChannelPricing(cp.ModelName, cp.ChannelId)
if existing != nil { if existing != nil {
existing.QuotaType = cp.QuotaType
existing.ModelRatio = cp.ModelRatio
existing.CompletionRatio = cp.CompletionRatio
existing.ModelPrice = cp.ModelPrice
existing.TagIds = cp.TagIds
existing.ApplyFields(cp.QuotaType, cp.ModelRatio, cp.CompletionRatio, cp.ModelPrice, cp.TagIds,
cp.CacheRatio, cp.CacheCreationRatio, cp.ImageRatio, cp.AudioRatio, cp.AudioCompletionRatio)
if err := existing.Update(); err != nil { if err := existing.Update(); err != nil {
common.ApiError(c, err) common.ApiError(c, err)
return return
@@ -168,6 +174,7 @@ func DeleteChannelPricing(c *gin.Context) {
return return
} }


common.SysLog(fmt.Sprintf("[ChannelPricing] deleted: id=%d", id))
common.ApiSuccess(c, nil) common.ApiSuccess(c, nil)
} }


@@ -204,6 +211,28 @@ func CopyGlobalPricing(c *gin.Context) {
continue continue
} }


// 获取全局扩展比率
globalCacheRatio, hasCacheRatio := ratio_setting.GetCacheRatio(ability.Model)
if !hasCacheRatio {
globalCacheRatio = 0
}
globalCacheCreationRatio, hasCacheCreationRatio := ratio_setting.GetCreateCacheRatio(ability.Model)
if !hasCacheCreationRatio {
globalCacheCreationRatio = 0
}
globalImageRatio, hasImageRatio := ratio_setting.GetImageRatio(ability.Model)
if !hasImageRatio {
globalImageRatio = 0
}
globalAudioRatio, hasAudioRatio := ratio_setting.GetAudioRatioV2(ability.Model)
if !hasAudioRatio {
globalAudioRatio = 0
}
globalAudioCompletionRatio, hasAudioCompRatio := ratio_setting.GetAudioCompletionRatioV2(ability.Model)
if !hasAudioCompRatio {
globalAudioCompletionRatio = 0
}

// 确定定价类型 // 确定定价类型
var quotaType int var quotaType int
var ratio, completionRatio, price float64 var ratio, completionRatio, price float64
@@ -224,28 +253,25 @@ func CopyGlobalPricing(c *gin.Context) {
} }


if existing != nil { if existing != nil {
existing.QuotaType = quotaType
existing.ModelRatio = ratio
existing.CompletionRatio = completionRatio
existing.ModelPrice = price
existing.ApplyFields(quotaType, ratio, completionRatio, price, "",
globalCacheRatio, globalCacheCreationRatio, globalImageRatio, globalAudioRatio, globalAudioCompletionRatio)
if err := existing.Update(); err == nil { if err := existing.Update(); err == nil {
imported++ imported++
} }
} else { } else {
cp := &model.ChannelPricing{ cp := &model.ChannelPricing{
ModelName: ability.Model,
ChannelId: channelId,
QuotaType: quotaType,
ModelRatio: ratio,
CompletionRatio: completionRatio,
ModelPrice: price,
ModelName: ability.Model,
ChannelId: channelId,
} }
cp.ApplyFields(quotaType, ratio, completionRatio, price, "",
globalCacheRatio, globalCacheCreationRatio, globalImageRatio, globalAudioRatio, globalAudioCompletionRatio)
if err := cp.Insert(); err == nil { if err := cp.Insert(); err == nil {
imported++ imported++
} }
} }
} }


common.SysLog(fmt.Sprintf("[ChannelPricing] copyGlobalPricing: channel=%d imported=%d/%d", channelId, imported, len(abilities)))
common.ApiSuccess(c, gin.H{ common.ApiSuccess(c, gin.H{
"total": len(abilities), "total": len(abilities),
"imported": imported, "imported": imported,
@@ -302,16 +328,7 @@ func GetChannelPricingWithTags(c *gin.Context) {
for _, cp := range list { for _, cp := range list {
item := &ChannelPricingWithTags{ item := &ChannelPricingWithTags{
ChannelPricing: cp, ChannelPricing: cp,
Tags: make([]*model.PricingTag, 0),
}
if cp.TagIds != "" {
for _, idStr := range strings.Split(cp.TagIds, ",") {
if id, err := strconv.Atoi(idStr); err == nil {
if tag, ok := tagMap[id]; ok {
item.Tags = append(item.Tags, tag)
}
}
}
Tags: model.ParseTagIds(cp.TagIds, tagMap),
} }
result = append(result, item) result = append(result, item)
} }
@@ -323,3 +340,51 @@ func GetChannelPricingWithTags(c *gin.Context) {
"items": result, "items": result,
}) })
} }

// SetDefaultChannelRequest 设置默认通道请求
type SetDefaultChannelRequest struct {
ModelName string `json:"model_name" binding:"required"`
ChannelId int `json:"channel_id" binding:"required"`
}

// SetDefaultChannel 设置指定模型的默认通道
func SetDefaultChannel(c *gin.Context) {
var req SetDefaultChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}

// 验证渠道定价记录存在
existing, err := model.GetChannelPricing(req.ModelName, req.ChannelId)
if err != nil || existing == nil {
common.ApiErrorMsg(c, "channel pricing not found for this model and channel")
return
}

if err := model.SetDefaultChannel(req.ModelName, req.ChannelId); err != nil {
common.ApiError(c, err)
return
}

common.SysLog(fmt.Sprintf("[ChannelPricing] set default: model=%s channel=%d", req.ModelName, req.ChannelId))
common.ApiSuccess(c, nil)
}

// ClearDefaultChannel 清除指定模型的默认通道
func ClearDefaultChannel(c *gin.Context) {
modelName := c.Param("name")
modelName = strings.TrimPrefix(modelName, "/")
if modelName == "" {
common.ApiErrorMsg(c, "model name is required")
return
}

if err := model.ClearDefaultChannel(modelName); err != nil {
common.ApiError(c, err)
return
}

common.SysLog(fmt.Sprintf("[ChannelPricing] cleared default: model=%s", modelName))
common.ApiSuccess(c, nil)
}

+ 111
- 0
controller/codex_channel_test.go Переглянути файл

@@ -0,0 +1,111 @@
package controller

import (
"bytes"
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func setupCodexChannelDB(t *testing.T, key string) *gorm.DB {
t.Helper()

db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)

sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
origLogDB := model.LOG_DB
model.DB = db
model.LOG_DB = db
common.UsingSQLite = true
common.RedisEnabled = false

require.NoError(t, db.AutoMigrate(&model.Channel{}))
require.NoError(t, db.Create(&model.Channel{
Id: 1,
Name: "codex-channel",
PublicName: "codex-channel",
Type: constant.ChannelTypeCodex,
Key: key,
Status: common.ChannelStatusEnabled,
}).Error)

t.Cleanup(func() {
model.DB = origDB
model.LOG_DB = origLogDB
_ = sqlDB.Close()
})

return db
}

func setupCodexChannelRouter(t *testing.T, key string) *gin.Engine {
t.Helper()
setupCodexChannelDB(t, key)

gin.SetMode(gin.TestMode)
r := gin.New()
g := r.Group("/api/channel")
g.POST("/:id/codex/refresh", RefreshCodexChannelCredential)
g.GET("/:id/codex/usage", GetCodexChannelUsage)
return r
}

func TestValidateChannelAcceptsCodexAPIKey(t *testing.T) {
channel := &model.Channel{
Name: "codex-api-key",
PublicName: "codex-api-key",
Type: constant.ChannelTypeCodex,
Key: "sk-codex-api-key",
}

require.NoError(t, validateChannel(channel, true))
}

func TestValidateChannelRejectsCodexOAuthWithoutAccountID(t *testing.T) {
channel := &model.Channel{
Name: "codex-oauth",
PublicName: "codex-oauth",
Type: constant.ChannelTypeCodex,
Key: `{"access_token":"token-only"}`,
}

err := validateChannel(channel, true)
require.Error(t, err)
assert.Contains(t, err.Error(), "account_id")
}

func TestRefreshCodexChannelCredentialRejectsAPIKeyMode(t *testing.T) {
router := setupCodexChannelRouter(t, "sk-codex-api-key")

req := httptest.NewRequest(http.MethodPost, "/api/channel/1/codex/refresh", bytes.NewReader([]byte(`{}`)))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
assert.Contains(t, w.Body.String(), "当前凭证方式不支持刷新凭证")
}

func TestGetCodexChannelUsageRejectsAPIKeyMode(t *testing.T) {
router := setupCodexChannelRouter(t, "sk-codex-api-key")

req := httptest.NewRequest(http.MethodGet, "/api/channel/1/codex/usage", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
assert.Contains(t, w.Body.String(), "当前凭证方式不支持查看用量")
}

+ 8
- 10
controller/codex_usage.go Переглянути файл

@@ -2,6 +2,7 @@ package controller


import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strconv" "strconv"
@@ -11,7 +12,6 @@ import (
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/relay/channel/codex"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"


"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -42,22 +42,19 @@ func GetCodexChannelUsage(c *gin.Context) {
return return
} }


oauthKey, err := codex.ParseOAuthKey(strings.TrimSpace(ch.Key))
oauthKey, err := common.ParseCodexOAuthCredential(strings.TrimSpace(ch.Key))
if err != nil { if err != nil {
if errors.Is(err, common.ErrCodexOAuthCredentialRequired) {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "当前凭证方式不支持查看用量"})
return
}
common.SysError("failed to parse oauth key: " + err.Error()) common.SysError("failed to parse oauth key: " + err.Error())
c.JSON(http.StatusOK, gin.H{"success": false, "message": "解析凭证失败,请检查渠道配置"}) c.JSON(http.StatusOK, gin.H{"success": false, "message": "解析凭证失败,请检查渠道配置"})
return return
} }

accessToken := strings.TrimSpace(oauthKey.AccessToken) accessToken := strings.TrimSpace(oauthKey.AccessToken)
accountID := strings.TrimSpace(oauthKey.AccountID) accountID := strings.TrimSpace(oauthKey.AccountID)
if accessToken == "" {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "codex channel: access_token is required"})
return
}
if accountID == "" {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "codex channel: account_id is required"})
return
}


client, err := service.NewProxyHttpClient(ch.GetSetting().Proxy) client, err := service.NewProxyHttpClient(ch.GetSetting().Proxy)
if err != nil { if err != nil {
@@ -98,6 +95,7 @@ func GetCodexChannelUsage(c *gin.Context) {


ctx2, cancel2 := context.WithTimeout(c.Request.Context(), 15*time.Second) ctx2, cancel2 := context.WithTimeout(c.Request.Context(), 15*time.Second)
defer cancel2() defer cancel2()

statusCode, body, err = service.FetchCodexWhamUsage(ctx2, client, ch.GetBaseURL(), oauthKey.AccessToken, accountID) statusCode, body, err = service.FetchCodexWhamUsage(ctx2, client, ch.GetBaseURL(), oauthKey.AccessToken, accountID)
if err != nil { if err != nil {
common.SysError("failed to fetch codex usage after refresh: " + err.Error()) common.SysError("failed to fetch codex usage after refresh: " + err.Error())


+ 128
- 0
controller/email_quota_rule.go Переглянути файл

@@ -0,0 +1,128 @@
package controller

import (
"strconv"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"

"github.com/gin-gonic/gin"
)

func GetAllEmailQuotaRules(c *gin.Context) {
list, err := model.GetAllEmailQuotaRules()
if err != nil {
common.ApiError(c, err)
return
}
common.ApiSuccess(c, list)
}

type CreateEmailQuotaRuleRequest struct {
EmailSuffix string `json:"email_suffix" binding:"required"`
Quota int64 `json:"quota" binding:"required"`
Enabled *bool `json:"enabled"`
Description string `json:"description"`
}

func CreateEmailQuotaRule(c *gin.Context) {
var req CreateEmailQuotaRuleRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}

existing, _ := model.GetEmailQuotaRuleBySuffix(req.EmailSuffix)
if existing != nil {
common.ApiErrorMsg(c, "该邮箱后缀已存在")
return
}

enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}

rule := &model.EmailQuotaRule{
EmailSuffix: req.EmailSuffix,
Quota: req.Quota,
Enabled: enabled,
Description: req.Description,
}
if err := rule.Insert(); err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, rule)
}

type UpdateEmailQuotaRuleRequest struct {
EmailSuffix string `json:"email_suffix"`
Quota *int64 `json:"quota"`
Enabled *bool `json:"enabled"`
Description *string `json:"description"`
}

func UpdateEmailQuotaRule(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.Atoi(idStr)
if err != nil {
common.ApiError(c, err)
return
}

var req UpdateEmailQuotaRuleRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}

rule, err := model.GetEmailQuotaRuleById(id)
if err != nil {
common.ApiError(c, err)
return
}

if req.EmailSuffix != "" && req.EmailSuffix != rule.EmailSuffix {
existing, _ := model.GetEmailQuotaRuleBySuffix(req.EmailSuffix)
if existing != nil && existing.Id != id {
common.ApiErrorMsg(c, "该邮箱后缀已存在")
return
}
rule.EmailSuffix = req.EmailSuffix
}
if req.Quota != nil {
rule.Quota = *req.Quota
}
if req.Enabled != nil {
rule.Enabled = *req.Enabled
}
if req.Description != nil {
rule.Description = *req.Description
}

if err := rule.Update(); err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, rule)
}

func DeleteEmailQuotaRule(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.Atoi(idStr)
if err != nil {
common.ApiError(c, err)
return
}

rule := &model.EmailQuotaRule{Id: id}
if err := rule.Delete(); err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, nil)
}

+ 264
- 0
controller/email_quota_rule_test.go Переглянути файл

@@ -0,0 +1,264 @@
package controller

import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/glebarez/sqlite"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func setupEmailQuotaRuleControllerDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
model.DB = db
common.UsingSQLite = true
common.RedisEnabled = false

require.NoError(t, db.AutoMigrate(&model.EmailQuotaRule{}))

t.Cleanup(func() {
model.DB = origDB
sqlDB.Close()
})
return db
}

func setupEmailQuotaRuleRouter() *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
g := r.Group("/api/email_quota_rule")
{
g.GET("/", GetAllEmailQuotaRules)
g.POST("/", CreateEmailQuotaRule)
g.PUT("/:id", UpdateEmailQuotaRule)
g.DELETE("/:id", DeleteEmailQuotaRule)
}
return r
}

func TestGetAllEmailQuotaRules_Empty(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/email_quota_rule/", nil)
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.True(t, resp["success"].(bool))
assert.Empty(t, resp["data"])
}

func TestCreateEmailQuotaRule_Success(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

body, _ := json.Marshal(map[string]interface{}{
"email_suffix": "@test.com",
"quota": 500000,
"description": "Test company",
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.True(t, resp["success"].(bool))

data := resp["data"].(map[string]interface{})
assert.Equal(t, "@test.com", data["email_suffix"])
assert.Equal(t, float64(500000), data["quota"])
assert.Equal(t, true, data["enabled"])
}

func TestCreateEmailQuotaRule_Duplicate(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

body, _ := json.Marshal(map[string]interface{}{
"email_suffix": "@dup.com",
"quota": 100,
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
var firstResp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &firstResp))
assert.True(t, firstResp["success"].(bool))

// Second create should fail
w2 := httptest.NewRecorder()
req2, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body))
req2.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w2, req2)

var resp2 map[string]interface{}
require.NoError(t, json.Unmarshal(w2.Body.Bytes(), &resp2))
assert.False(t, resp2["success"].(bool))
}

func TestCreateEmailQuotaRule_MissingFields(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

body, _ := json.Marshal(map[string]interface{}{
"description": "no suffix",
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.False(t, resp["success"].(bool))
}

func TestUpdateEmailQuotaRule_Success(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

// Create first
rule := &model.EmailQuotaRule{EmailSuffix: "@up.com", Quota: 100, Enabled: true}
require.NoError(t, rule.Insert())

body, _ := json.Marshal(map[string]interface{}{
"quota": 999,
"description": "updated desc",
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/email_quota_rule/"+itoa(rule.Id), bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.True(t, resp["success"].(bool))

data := resp["data"].(map[string]interface{})
assert.Equal(t, float64(999), data["quota"])
assert.Equal(t, "updated desc", data["description"])
}

func TestUpdateEmailQuotaRule_ToggleEnabled(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

rule := &model.EmailQuotaRule{EmailSuffix: "@toggle.com", Quota: 500, Enabled: true}
require.NoError(t, rule.Insert())

body, _ := json.Marshal(map[string]interface{}{
"enabled": false,
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/email_quota_rule/"+itoa(rule.Id), bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.True(t, resp["success"].(bool))

data := resp["data"].(map[string]interface{})
assert.Equal(t, false, data["enabled"])

// Cache should reflect disabled
assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@toggle.com"))
}

func TestDeleteEmailQuotaRule_Success(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

rule := &model.EmailQuotaRule{EmailSuffix: "@del.com", Quota: 100, Enabled: true}
require.NoError(t, rule.Insert())

w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/email_quota_rule/"+itoa(rule.Id), nil)
router.ServeHTTP(w, req)

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.True(t, resp["success"].(bool))

// Verify deleted
assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@del.com"))
}

func TestCRUD_FullFlow(t *testing.T) {
setupEmailQuotaRuleControllerDB(t)
router := setupEmailQuotaRuleRouter()

// 1. Create
body, _ := json.Marshal(map[string]interface{}{
"email_suffix": "@full.com",
"quota": 1000,
"description": "full flow test",
})
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
var createResp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &createResp))
assert.True(t, createResp["success"].(bool))

// 2. List
w = httptest.NewRecorder()
req, _ = http.NewRequest("GET", "/api/email_quota_rule/", nil)
router.ServeHTTP(w, req)
var listResp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &listResp))
data := listResp["data"].([]interface{})
assert.Len(t, data, 1)

// 3. Update
ruleId := itoa(int(data[0].(map[string]interface{})["id"].(float64)))
body, _ = json.Marshal(map[string]interface{}{"quota": 2000})
w = httptest.NewRecorder()
req, _ = http.NewRequest("PUT", "/api/email_quota_rule/"+ruleId, bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
var updateResp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &updateResp))
assert.True(t, updateResp["success"].(bool))

// 4. Verify cache
assert.Equal(t, int64(2000), model.MatchEmailQuotaRule("user@full.com"))

// 5. Delete
w = httptest.NewRecorder()
req, _ = http.NewRequest("DELETE", "/api/email_quota_rule/"+ruleId, nil)
router.ServeHTTP(w, req)
var delResp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &delResp))
assert.True(t, delResp["success"].(bool))

// 6. Verify cache cleared
assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@full.com"))
}

func itoa(i int) string {
return strconv.Itoa(i)
}

+ 74
- 5
controller/misc.go Переглянути файл

@@ -61,12 +61,13 @@ func GetStatus(c *gin.Context) {
"telegram_oauth": common.TelegramOAuthEnabled, "telegram_oauth": common.TelegramOAuthEnabled,
"telegram_bot_name": common.TelegramBotName, "telegram_bot_name": common.TelegramBotName,
"system_name": common.SystemName, "system_name": common.SystemName,
"logo": common.Logo,
"logo": common.GetEffectiveLogo(),
"footer_html": common.Footer, "footer_html": common.Footer,
"wechat_qrcode": common.WeChatAccountQRCodeImageURL, "wechat_qrcode": common.WeChatAccountQRCodeImageURL,
"wechat_login": common.WeChatAuthEnabled, "wechat_login": common.WeChatAuthEnabled,
"server_address": system_setting.ServerAddress, "server_address": system_setting.ServerAddress,
"turnstile_check": common.TurnstileCheckEnabled, "turnstile_check": common.TurnstileCheckEnabled,
"captcha_enabled": common.CaptchaEnabled,
"turnstile_site_key": common.TurnstileSiteKey, "turnstile_site_key": common.TurnstileSiteKey,
"top_up_link": common.TopUpLink, "top_up_link": common.TopUpLink,
"docs_link": operation_setting.GetGeneralSetting().DocsLink, "docs_link": operation_setting.GetGeneralSetting().DocsLink,
@@ -87,6 +88,7 @@ func GetStatus(c *gin.Context) {
"demo_site_enabled": operation_setting.DemoSiteEnabled, "demo_site_enabled": operation_setting.DemoSiteEnabled,
"self_use_mode_enabled": operation_setting.SelfUseModeEnabled, "self_use_mode_enabled": operation_setting.SelfUseModeEnabled,
"default_use_auto_group": setting.DefaultUseAutoGroup, "default_use_auto_group": setting.DefaultUseAutoGroup,
"default_language": common.DefaultLanguage,


"usd_exchange_rate": operation_setting.USDExchangeRate, "usd_exchange_rate": operation_setting.USDExchangeRate,
"price": operation_setting.Price, "price": operation_setting.Price,
@@ -113,8 +115,10 @@ func GetStatus(c *gin.Context) {
"passkey_user_verification": passkeySetting.UserVerification, "passkey_user_verification": passkeySetting.UserVerification,
"passkey_attachment": passkeySetting.AttachmentPreference, "passkey_attachment": passkeySetting.AttachmentPreference,
"setup": constant.Setup, "setup": constant.Setup,
"user_agreement_enabled": legalSetting.UserAgreement != "",
"privacy_policy_enabled": legalSetting.PrivacyPolicy != "",
"user_agreement_enabled": legalSetting.UserAgreementZh != "" || legalSetting.UserAgreementEn != "",
"privacy_policy_enabled": legalSetting.PrivacyPolicyZh != "" || legalSetting.PrivacyPolicyEn != "",
"terms_enabled": legalSetting.TermsOfServiceZh != "" || legalSetting.TermsOfServiceEn != "",
"usage_policy_enabled": legalSetting.UsagePolicyZh != "" || legalSetting.UsagePolicyEn != "",
"checkin_enabled": operation_setting.GetCheckinSetting().Enabled, "checkin_enabled": operation_setting.GetCheckinSetting().Enabled,
"_qn": "new-api", "_qn": "new-api",
} }
@@ -188,20 +192,53 @@ func GetAbout(c *gin.Context) {
return return
} }


func getLegalContent(zh, en string, lang string) string {
if lang == "en" {
return en
}
return zh
}

func GetUserAgreement(c *gin.Context) { func GetUserAgreement(c *gin.Context) {
ls := system_setting.GetLegalSettings()
lang := c.Query("lang")
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
"message": "", "message": "",
"data": system_setting.GetLegalSettings().UserAgreement,
"data": getLegalContent(ls.UserAgreementZh, ls.UserAgreementEn, lang),
}) })
return return
} }


func GetPrivacyPolicy(c *gin.Context) { func GetPrivacyPolicy(c *gin.Context) {
ls := system_setting.GetLegalSettings()
lang := c.Query("lang")
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": getLegalContent(ls.PrivacyPolicyZh, ls.PrivacyPolicyEn, lang),
})
return
}

func GetTermsOfService(c *gin.Context) {
ls := system_setting.GetLegalSettings()
lang := c.Query("lang")
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": getLegalContent(ls.TermsOfServiceZh, ls.TermsOfServiceEn, lang),
})
return
}

func GetUsagePolicy(c *gin.Context) {
ls := system_setting.GetLegalSettings()
lang := c.Query("lang")
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
"message": "", "message": "",
"data": system_setting.GetLegalSettings().PrivacyPolicy,
"data": getLegalContent(ls.UsagePolicyZh, ls.UsagePolicyEn, lang),
}) })
return return
} }
@@ -228,7 +265,39 @@ func GetHomePageContent(c *gin.Context) {
return return
} }


func GetCaptcha(c *gin.Context) {
if !common.CaptchaEnabled {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "验证码功能未启用",
})
return
}
id, b64s, err := common.GenerateCaptcha()
if err != nil {
common.ApiErrorMsg(c, "生成验证码失败")
return
}
common.ApiSuccess(c, gin.H{
"id": id,
"captcha_image": b64s,
})
}

func SendEmailVerification(c *gin.Context) { func SendEmailVerification(c *gin.Context) {
if common.CaptchaEnabled {
captchaId := c.Query("captcha_id")
captchaCode := c.Query("captcha_code")
if captchaId == "" || captchaCode == "" {
common.ApiErrorMsg(c, "请先完成图片验证码")
return
}
if !common.VerifyCaptcha(captchaId, captchaCode) {
common.ApiErrorMsg(c, "图片验证码错误或已过期")
return
}
}

email := c.Query("email") email := c.Query("email")
if err := common.Validate.Var(email, "required,email"); err != nil { if err := common.Validate.Var(email, "required,email"); err != nil {
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{


+ 62
- 0
controller/playground_channels.go Переглянути файл

@@ -0,0 +1,62 @@
package controller

import (
"net/http"

"github.com/QuantumNous/new-api/model"
"github.com/gin-gonic/gin"
)

func GetModelChannels(c *gin.Context) {
modelName := c.Query("model")
if modelName == "" {
c.JSON(http.StatusBadRequest, gin.H{
"success": false,
"message": "model parameter is required",
})
return
}

userId := c.GetInt("id")
if userId == 0 {
c.JSON(http.StatusUnauthorized, gin.H{
"success": false,
"message": "unauthorized",
})
return
}

userCache, err := model.GetUserCache(userId)
if err != nil || userCache == nil {
c.JSON(http.StatusInternalServerError, gin.H{
"success": false,
"message": "failed to get user info",
})
return
}
userGroup := userCache.Group
if userGroup == "" {
c.JSON(http.StatusInternalServerError, gin.H{
"success": false,
"message": "failed to get user group",
})
return
}

channels, defaultChannelId, err := model.GetModelChannelsForGroup(modelName, userGroup)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"success": false,
"message": "failed to query channels",
})
return
}

c.JSON(http.StatusOK, gin.H{
"success": true,
"data": gin.H{
"channels": channels,
"default_channel_id": defaultChannelId,
},
})
}

+ 2
- 0
controller/redemption.go Переглянути файл

@@ -87,6 +87,7 @@ func AddRedemption(c *gin.Context) {
cleanRedemption := model.Redemption{ cleanRedemption := model.Redemption{
UserId: c.GetInt("id"), UserId: c.GetInt("id"),
Name: redemption.Name, Name: redemption.Name,
Remark: redemption.Remark,
Key: key, Key: key,
CreatedTime: common.GetTimestamp(), CreatedTime: common.GetTimestamp(),
Quota: redemption.Quota, Quota: redemption.Quota,
@@ -146,6 +147,7 @@ func UpdateRedemption(c *gin.Context) {
} }
// If you add more fields, please also update redemption.Update() // If you add more fields, please also update redemption.Update()
cleanRedemption.Name = redemption.Name cleanRedemption.Name = redemption.Name
cleanRedemption.Remark = redemption.Remark
cleanRedemption.Quota = redemption.Quota cleanRedemption.Quota = redemption.Quota
cleanRedemption.ExpiredTime = redemption.ExpiredTime cleanRedemption.ExpiredTime = redemption.ExpiredTime
} }


+ 159
- 0
controller/redemption_test.go Переглянути файл

@@ -0,0 +1,159 @@
package controller

import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/glebarez/sqlite"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func setupRedemptionControllerDB(t *testing.T) *gorm.DB {
t.Helper()

db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)

sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
origLogDB := model.LOG_DB
model.DB = db
model.LOG_DB = db
common.UsingSQLite = true
common.RedisEnabled = false

require.NoError(t, db.AutoMigrate(&model.User{}, &model.Redemption{}, &model.Log{}))

t.Cleanup(func() {
model.DB = origDB
model.LOG_DB = origLogDB
_ = sqlDB.Close()
})

return db
}

func setupRedemptionRouter() *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(func(c *gin.Context) {
c.Set("id", 1)
c.Next()
})

g := r.Group("/api/redemption")
g.POST("/", AddRedemption)
g.PUT("/", UpdateRedemption)

return r
}

func TestAddRedemptionStoresRemark(t *testing.T) {
db := setupRedemptionControllerDB(t)
router := setupRedemptionRouter()

body, err := json.Marshal(map[string]interface{}{
"name": "campaign-a",
"remark": "admin only note",
"quota": 500000,
"count": 1,
"expired_time": 0,
})
require.NoError(t, err)

req := httptest.NewRequest(http.MethodPost, "/api/redemption/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)

var rows []model.Redemption
require.NoError(t, db.Find(&rows).Error)
require.Len(t, rows, 1)
assert.Equal(t, "campaign-a", rows[0].Name)
assert.Equal(t, "admin only note", rows[0].Remark)
}

func TestUpdateRedemptionStoresRemark(t *testing.T) {
db := setupRedemptionControllerDB(t)
router := setupRedemptionRouter()

row := model.Redemption{
Id: 1,
UserId: 1,
Key: "update-remark-key",
Name: "campaign-b",
Remark: "before update",
Status: common.RedemptionCodeStatusEnabled,
Quota: 500000,
CreatedTime: common.GetTimestamp(),
}
require.NoError(t, db.Create(&row).Error)

body, err := json.Marshal(map[string]interface{}{
"id": 1,
"name": "campaign-b",
"remark": "after update",
"quota": 500000,
"expired_time": 0,
})
require.NoError(t, err)

req := httptest.NewRequest(http.MethodPut, "/api/redemption/", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)

var stored model.Redemption
require.NoError(t, db.First(&stored, 1).Error)
assert.Equal(t, "after update", stored.Remark)
}

func TestUpdateRedemptionStatusOnlyKeepsRemark(t *testing.T) {
db := setupRedemptionControllerDB(t)
router := setupRedemptionRouter()

row := model.Redemption{
Id: 1,
UserId: 1,
Key: "status-only-remark-key",
Name: "campaign-c",
Remark: "keep this remark",
Status: common.RedemptionCodeStatusEnabled,
Quota: 500000,
CreatedTime: common.GetTimestamp(),
}
require.NoError(t, db.Create(&row).Error)

body, err := json.Marshal(map[string]interface{}{
"id": 1,
"status": common.RedemptionCodeStatusDisabled,
"remark": "should not overwrite",
})
require.NoError(t, err)

req := httptest.NewRequest(http.MethodPut, "/api/redemption/?status_only=true", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)

var stored model.Redemption
require.NoError(t, db.First(&stored, 1).Error)
assert.Equal(t, common.RedemptionCodeStatusDisabled, stored.Status)
assert.Equal(t, "keep this remark", stored.Remark)
}

+ 6
- 0
controller/relay.go Переглянути файл

@@ -369,6 +369,12 @@ func processChannelError(c *gin.Context, channelError types.ChannelError, err *t
other["channel_id"] = channelId other["channel_id"] = channelId
other["channel_name"] = c.GetString("channel_name") other["channel_name"] = c.GetString("channel_name")
other["channel_type"] = c.GetInt("channel_type") other["channel_type"] = c.GetInt("channel_type")
if err.UpstreamRequestId != "" {
other["upstream_request_id"] = err.UpstreamRequestId
}
if err.UpstreamBody != "" {
other["upstream_body"] = err.UpstreamBody
}
adminInfo := make(map[string]interface{}) adminInfo := make(map[string]interface{})
adminInfo["use_channel"] = c.GetStringSlice("use_channel") adminInfo["use_channel"] = c.GetStringSlice("use_channel")
isMultiKey := common.GetContextKeyBool(c, constant.ContextKeyChannelIsMultiKey) isMultiKey := common.GetContextKeyBool(c, constant.ContextKeyChannelIsMultiKey)


+ 55
- 0
controller/reorder.go Переглянути файл

@@ -0,0 +1,55 @@
package controller

import (
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"

"github.com/gin-gonic/gin"
)

// ReorderVendors 批量更新供应商排序
func ReorderVendors(c *gin.Context) {
var req struct {
Items []struct {
Id int `json:"id"`
SortOrder int `json:"sort_order"`
} `json:"items"`
}
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}
if len(req.Items) == 0 {
common.ApiErrorMsg(c, "items 不能为空")
return
}
if err := model.ReorderVendors(req.Items); err != nil {
common.ApiError(c, err)
return
}
common.ApiSuccess(c, nil)
}

// ReorderModels 批量更新模型排序
func ReorderModels(c *gin.Context) {
var req struct {
Items []struct {
Id int `json:"id"`
SortOrder int `json:"sort_order"`
} `json:"items"`
}
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}
if len(req.Items) == 0 {
common.ApiErrorMsg(c, "items 不能为空")
return
}
if err := model.ReorderModels(req.Items); err != nil {
common.ApiError(c, err)
return
}
model.RefreshPricing()
common.ApiSuccess(c, nil)
}

+ 47
- 4
controller/topup.go Переглянути файл

@@ -72,17 +72,38 @@ func GetTopUpInfo(c *gin.Context) {
payMethods = append(payMethods, wechatMethod) payMethods = append(payMethods, wechatMethod)
} }
} }
// 如果启用了支付宝支付,添加到支付方法列表
if setting.IsAlipayConfigured() {
hasAlipay := false
for _, method := range payMethods {
if method["type"] == PaymentMethodAlipay {
hasAlipay = true
break
}
}
if !hasAlipay {
alipayMethod := map[string]string{
"name": "Alipay",
"type": PaymentMethodAlipay,
"color": "rgba(var(--semi-blue-5), 1)",
"min_topup": strconv.Itoa(setting.AlipayMinTopUp),
}
payMethods = append(payMethods, alipayMethod)
}
}


data := gin.H{ data := gin.H{
"enable_online_topup": enableOnlineTopup, "enable_online_topup": enableOnlineTopup,
"enable_stripe_topup": setting.StripeApiSecret != "" && setting.StripeWebhookSecret != "" && setting.StripePriceId != "", "enable_stripe_topup": setting.StripeApiSecret != "" && setting.StripeWebhookSecret != "" && setting.StripePriceId != "",
"enable_creem_topup": setting.CreemApiKey != "" && setting.CreemProducts != "[]", "enable_creem_topup": setting.CreemApiKey != "" && setting.CreemProducts != "[]",
"enable_wechat_topup": setting.IsWechatPayConfigured(), "enable_wechat_topup": setting.IsWechatPayConfigured(),
"enable_alipay_topup": setting.IsAlipayConfigured(),
"creem_products": setting.CreemProducts, "creem_products": setting.CreemProducts,
"pay_methods": payMethods, "pay_methods": payMethods,
"min_topup": operation_setting.MinTopUp, "min_topup": operation_setting.MinTopUp,
"stripe_min_topup": setting.StripeMinTopUp, "stripe_min_topup": setting.StripeMinTopUp,
"wechat_pay_min_topup": setting.WechatPayMinTopUp, "wechat_pay_min_topup": setting.WechatPayMinTopUp,
"alipay_pay_min_topup": setting.AlipayMinTopUp,
"amount_options": operation_setting.GetPaymentSetting().AmountOptions, "amount_options": operation_setting.GetPaymentSetting().AmountOptions,
"discount": operation_setting.GetPaymentSetting().AmountDiscount, "discount": operation_setting.GetPaymentSetting().AmountDiscount,
} }
@@ -143,15 +164,37 @@ func getPayMoney(amount int64, group string) float64 {
} }


func getMinTopup() int64 { func getMinTopup() int64 {
minTopup := operation_setting.MinTopUp
return calcMinTopup(operation_setting.MinTopUp)
}

// calcMinTopup 计算最低充值数量(考虑 QuotaDisplayType 换算)
func calcMinTopup(baseMinTopup int) int64 {
minTopup := baseMinTopup
if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens { if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens {
dMinTopup := decimal.NewFromInt(int64(minTopup))
dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit)
minTopup = int(dMinTopup.Mul(dQuotaPerUnit).IntPart())
minTopup = minTopup * int(common.QuotaPerUnit)
} }
return int64(minTopup) return int64(minTopup)
} }


// calcPayMoney 计算应付金额(元),使用指定的单价和最低充值
func calcPayMoney(amount float64, group string, unitPrice float64) float64 {
originalAmount := amount
if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens {
amount = amount / common.QuotaPerUnit
}
topupGroupRatio := common.GetTopupGroupRatio(group)
if topupGroupRatio == 0 {
topupGroupRatio = 1
}
discount := 1.0
if ds, ok := operation_setting.GetPaymentSetting().AmountDiscount[int(originalAmount)]; ok {
if ds > 0 {
discount = ds
}
}
return amount * unitPrice * topupGroupRatio * discount
}

func RequestEpay(c *gin.Context) { func RequestEpay(c *gin.Context) {
var req EpayRequest var req EpayRequest
err := c.ShouldBindJSON(&req) err := c.ShouldBindJSON(&req)


+ 280
- 0
controller/topup_alipay.go Переглянути файл

@@ -0,0 +1,280 @@
package controller

import (
"context"
"encoding/base64"
"fmt"
"log"
"net/http"
"strconv"
"sync"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting"

"github.com/gin-gonic/gin"
"github.com/go-pay/gopay"
"github.com/go-pay/gopay/alipay"
"github.com/thanhpk/randstr"
)

const (
PaymentMethodAlipay = "alipay"
)

var alipayClientMu sync.Mutex
var alipayClient *alipay.Client

// ResetAlipayClient 重置支付宝客户端(配置变更时调用)
func ResetAlipayClient() {
alipayClientMu.Lock()
alipayClient = nil
alipayClientMu.Unlock()
}

func init() {
setting.OnAlipayConfigChanged = ResetAlipayClient
}

// getAlipayClient 获取或创建支付宝客户端
func getAlipayClient() (*alipay.Client, error) {
alipayClientMu.Lock()
defer alipayClientMu.Unlock()

if alipayClient != nil {
return alipayClient, nil
}

if !setting.IsAlipayConfigured() {
return nil, fmt.Errorf("支付宝未配置")
}

client, err := alipay.NewClient(setting.AlipayAppID, setting.AlipayPrivateKey, true)
if err != nil {
return nil, fmt.Errorf("创建支付宝客户端失败: %w", err)
}

client.SetCharset("utf-8").
SetSignType(alipay.RSA2).
SetNotifyUrl(setting.AlipayNotifyURL)

// 设置支付宝公钥(用于回调验签)
pubKeyBytes, err := base64.StdEncoding.DecodeString(setting.AlipayPublicKey)
if err != nil {
pubKeyBytes = []byte(setting.AlipayPublicKey)
}
client.AutoVerifySign([]byte(wrapAsPEM(pubKeyBytes, "PUBLIC KEY")))

alipayClient = client
return alipayClient, nil
}

// AlipayPayRequest 支付宝支付请求参数
type AlipayPayRequest struct {
Amount int64 `json:"amount"`
}

// createAlipayPrecreateOrder 调用支付宝当面付预下单 API,返回二维码内容
func createAlipayPrecreateOrder(client *alipay.Client, subject, tradeNo string, totalAmount string) (string, error) {
bm := make(gopay.BodyMap)
bm.Set("subject", subject).
Set("out_trade_no", tradeNo).
Set("total_amount", totalAmount).
Set("product_code", "FACE_TO_FACE_PAYMENT")

ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()

rsp, err := client.TradePrecreate(ctx, bm)
if err != nil {
return "", fmt.Errorf("支付宝当面付下单失败: %w", err)
}

if rsp.Response.Code != "10000" {
return "", fmt.Errorf("支付宝错误: %s - %s", rsp.Response.Code, rsp.Response.Msg)
}

return rsp.Response.QrCode, nil
}

// RequestAlipayPayAmount 计算支付宝应付金额
func RequestAlipayPayAmount(c *gin.Context) {
var req AlipayPayRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(200, gin.H{"message": "error", "data": "参数错误"})
return
}

minTopup := calcMinTopup(setting.AlipayMinTopUp)
if req.Amount < minTopup {
c.JSON(200, gin.H{"message": "error", "data": fmt.Sprintf("充值数量不能小于 %d", minTopup)})
return
}

id := c.GetInt("id")
group, err := model.GetUserGroup(id, true)
if err != nil {
c.JSON(200, gin.H{"message": "error", "data": "获取用户分组失败"})
return
}

payMoney := calcPayMoney(float64(req.Amount), group, setting.AlipayUnitPrice)
if payMoney <= 0.01 {
c.JSON(200, gin.H{"message": "error", "data": "充值金额过低"})
return
}

c.JSON(200, gin.H{"message": "success", "data": strconv.FormatFloat(payMoney, 'f', 2, 64)})
}

// RequestAlipayPay 创建支付宝支付订单,返回二维码 URL
func RequestAlipayPay(c *gin.Context) {
var req AlipayPayRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(200, gin.H{"message": "error", "data": "参数错误"})
return
}

minTopup := calcMinTopup(setting.AlipayMinTopUp)
if req.Amount < minTopup {
c.JSON(200, gin.H{"message": "error", "data": fmt.Sprintf("充值数量不能小于 %d", minTopup)})
return
}
if req.Amount > 10000 {
c.JSON(200, gin.H{"message": "error", "data": "充值数量不能大于 10000"})
return
}

if !setting.IsAlipayConfigured() {
c.JSON(200, gin.H{"message": "error", "data": "支付宝未配置"})
return
}

client, err := getAlipayClient()
if err != nil {
log.Println("获取支付宝客户端失败:", err)
c.JSON(200, gin.H{"message": "error", "data": "支付宝配置错误"})
return
}

id := c.GetInt("id")

group, err := model.GetUserGroup(id, true)
if err != nil {
c.JSON(200, gin.H{"message": "error", "data": "获取用户分组失败"})
return
}

topupGroupRatio := common.GetTopupGroupRatio(group)
if topupGroupRatio == 0 {
topupGroupRatio = 1
}
chargedMoney := float64(req.Amount) * topupGroupRatio

tradeNo := fmt.Sprintf("ali%d%s", time.Now().UnixMilli(), randstr.String(8))

payMoney := calcPayMoney(float64(req.Amount), group, setting.AlipayUnitPrice)
totalAmount := strconv.FormatFloat(payMoney, 'f', 2, 64)

qrCode, err := createAlipayPrecreateOrder(client, fmt.Sprintf("充值%d", req.Amount), tradeNo, totalAmount)
if err != nil {
log.Println(err)
c.JSON(200, gin.H{"message": "error", "data": "拉起支付失败"})
return
}

topUp := &model.TopUp{
UserId: id,
Amount: req.Amount,
Money: chargedMoney,
TradeNo: tradeNo,
PaymentMethod: PaymentMethodAlipay,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
}
if err := topUp.Insert(); err != nil {
c.JSON(200, gin.H{"message": "error", "data": "创建订单失败"})
return
}

c.JSON(200, gin.H{
"message": "success",
"data": gin.H{
"trade_no": tradeNo,
"qr_code_url": qrCode,
},
})
}

// AlipayPayStatus 轮询支付宝支付订单状态
func AlipayPayStatus(c *gin.Context) {
tradeNo := c.Query("trade_no")
if tradeNo == "" {
c.JSON(200, gin.H{"message": "error", "data": "参数错误"})
return
}

topUp := model.GetTopUpByTradeNo(tradeNo)
if topUp == nil {
c.JSON(200, gin.H{"message": "error", "data": "订单不存在"})
return
}

userId := c.GetInt("id")
if topUp.UserId != userId {
c.JSON(200, gin.H{"message": "error", "data": "订单不存在"})
return
}

c.JSON(200, gin.H{
"message": "success",
"data": gin.H{
"status": topUp.Status,
"amount": topUp.Amount,
},
})
}

// AlipayPayWebhook 处理支付宝异步回调通知
func AlipayPayWebhook(c *gin.Context) {
notifyReq, err := alipay.ParseNotifyToBodyMap(c.Request)
if err != nil {
log.Printf("解析支付宝回调失败: %v", err)
c.String(http.StatusBadRequest, "fail")
return
}

ok, err := alipay.VerifySign(setting.AlipayPublicKey, notifyReq)
if err != nil {
log.Printf("支付宝回调验签失败: %v", err)
c.String(http.StatusBadRequest, "fail")
return
}
if !ok {
log.Printf("支付宝回调验签不通过")
c.String(http.StatusBadRequest, "fail")
return
}

tradeStatus := notifyReq.Get("trade_status")
if tradeStatus != "TRADE_SUCCESS" {
c.String(http.StatusOK, "success")
return
}

tradeNo := notifyReq.Get("out_trade_no")

LockOrder(tradeNo)
defer UnlockOrder(tradeNo)

if err := model.RechargeAlipay(tradeNo); err != nil {
log.Printf("支付宝充值失败: %s, err: %s", tradeNo, err.Error())
c.String(http.StatusInternalServerError, "fail")
return
}

log.Printf("支付宝充值成功: %s", tradeNo)
c.String(http.StatusOK, "success")
}

+ 2
- 24
controller/topup_wechat.go Переглянути файл

@@ -18,7 +18,6 @@ import (
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting" "github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/operation_setting"


"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/go-pay/gopay" "github.com/go-pay/gopay"
@@ -355,31 +354,10 @@ func WechatPayWebhook(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": "SUCCESS", "message": "成功"}) c.JSON(http.StatusOK, gin.H{"code": "SUCCESS", "message": "成功"})
} }


// getWechatPayMoney 计算微信支付应付金额(元)
func getWechatPayMoney(amount float64, group string) float64 { func getWechatPayMoney(amount float64, group string) float64 {
originalAmount := amount
if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens {
amount = amount / common.QuotaPerUnit
}
topupGroupRatio := common.GetTopupGroupRatio(group)
if topupGroupRatio == 0 {
topupGroupRatio = 1
}
discount := 1.0
if ds, ok := operation_setting.GetPaymentSetting().AmountDiscount[int(originalAmount)]; ok {
if ds > 0 {
discount = ds
}
}
payMoney := amount * setting.WechatPayUnitPrice * topupGroupRatio * discount
return payMoney
return calcPayMoney(amount, group, setting.WechatPayUnitPrice)
} }


// getWechatMinTopup 获取微信支付最低充值数量
func getWechatMinTopup() int64 { func getWechatMinTopup() int64 {
minTopup := setting.WechatPayMinTopUp
if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens {
minTopup = minTopup * int(common.QuotaPerUnit)
}
return int64(minTopup)
return calcMinTopup(setting.WechatPayMinTopUp)
} }

+ 5
- 1
controller/user.go Переглянути файл

@@ -270,6 +270,7 @@ func GetUser(c *gin.Context) {
common.ApiErrorI18n(c, i18n.MsgUserNoPermissionSameLevel) common.ApiErrorI18n(c, i18n.MsgUserNoPermissionSameLevel)
return return
} }
user.ApplySyncedQuota()
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
"message": "", "message": "",
@@ -523,13 +524,16 @@ func GetUserModels(c *gin.Context) {
} }
groups := service.GetUserUsableGroups(user.Group) groups := service.GetUserUsableGroups(user.Group)
var models []string var models []string
seen := make(map[string]struct{})
for group := range groups { for group := range groups {
for _, g := range model.GetGroupEnabledModels(group) { for _, g := range model.GetGroupEnabledModels(group) {
if !common.StringsContains(models, g) {
if _, ok := seen[g]; !ok {
seen[g] = struct{}{}
models = append(models, g) models = append(models, g)
} }
} }
} }
models = common.StringsSubtract(models, model.GetDisabledModelNames(models))
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
"message": "", "message": "",


+ 1
- 0
docs/DATABASE_SCHEMA.md Переглянути файл

@@ -189,6 +189,7 @@
| key | string | 兑换码(32字符,唯一) | | key | string | 兑换码(32字符,唯一) |
| status | int | 状态:1=启用,2=已使用,3=已禁用 | | status | int | 状态:1=启用,2=已使用,3=已禁用 |
| name | string | 兑换码名称 | | name | string | 兑换码名称 |
| remark | string | 备注 |
| quota | int | 额度值 | | quota | int | 额度值 |
| created_time | int64 | 创建时间 | | created_time | int64 | 创建时间 |
| redeemed_time | int64 | 兑换时间 | | redeemed_time | int64 | 兑换时间 |


+ 307
- 0
docs/superpowers/plans/2026-04-17-default-language-setting.md Переглянути файл

@@ -0,0 +1,307 @@
# Default Language Setting Implementation Plan

> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.

**Goal:** Allow admin to configure a global default language in System Settings, which is forced for unauthenticated users and logged-in users without a personal language preference.

**Architecture:** Backend adds a `DefaultLanguage` option to the existing flat-key option system (same pattern as `SystemName`, `ServerAddress`). The value flows to frontend via `/api/status`. Frontend applies it in two places: `PageLayout.jsx` (for unauthenticated users) and `UserContext` (for logged-in users without preference).

**Tech Stack:** Go (backend), React + Semi UI (frontend), i18next (i18n)

**Spec:** `docs/superpowers/specs/2026-04-17-default-language-setting-design.md`

---

## File Map

| File | Action | Responsibility |
|------|--------|---------------|
| `common/constants.go` | Modify | Declare `DefaultLanguage` variable |
| `model/option.go` | Modify | Handle `DefaultLanguage` in option update switch |
| `controller/misc.go` | Modify | Expose `default_language` in `/api/status` response |
| `web/src/components/settings/SystemSetting.jsx` | Modify | Admin UI: language dropdown in system settings |
| `web/src/components/layout/PageLayout.jsx` | Modify | Apply default language for unauthenticated users |
| `web/src/context/User/index.jsx` | Modify | Fall back to default language when user has no preference |

---

### Task 1: Backend — Add DefaultLanguage variable and option handling

**Files:**
- Modify: `common/constants.go:17-18` (after `TopUpLink` declaration)
- Modify: `model/option.go:457-458` (after `Logo` case in `updateOptionMap` switch)

- [ ] **Step 1: Add `DefaultLanguage` variable to `common/constants.go`**

In `common/constants.go`, after line 18 (`var TopUpLink = ""`), add:

```go
var DefaultLanguage = "" // admin-configured default language; empty = follow browser detection
```

- [ ] **Step 2: Add case handler in `model/option.go`**

In `model/option.go`, after the `case "Logo":` block (line 457-458), add:

```go
case "DefaultLanguage":
common.DefaultLanguage = value
```

- [ ] **Step 3: Verify the Go code compiles**

Run: `go build ./...`
Expected: no errors

- [ ] **Step 4: Commit**

```bash
git add common/constants.go model/option.go
git commit -m "feat(default-language): add DefaultLanguage variable and option handler"
```

---

### Task 2: Backend — Expose default_language in /api/status

**Files:**
- Modify: `controller/misc.go:89` (after `"default_use_auto_group"` line)

- [ ] **Step 1: Add `default_language` field to GetStatus response**

In `controller/misc.go`, inside `GetStatus()`, after the line containing `"default_use_auto_group": setting.DefaultUseAutoGroup,` (line 89), add:

```go
"default_language": common.DefaultLanguage,
```

Place it right before the blank line at line 90, aligned with the surrounding entries.

- [ ] **Step 2: Verify the Go code compiles**

Run: `go build ./...`
Expected: no errors

- [ ] **Step 3: Commit**

```bash
git add controller/misc.go
git commit -m "feat(default-language): expose default_language in /api/status response"
```

---

### Task 3: Frontend — Add DefaultLanguage dropdown in System Settings

**Files:**
- Modify: `web/src/components/settings/SystemSetting.jsx`

- [ ] **Step 1: Add `DefaultLanguage` to `inputs` state**

In `web/src/components/settings/SystemSetting.jsx`, in the `inputs` useState (around line 49-112), add `DefaultLanguage: ''` after `ServerAddress: ''` (line 102):

```javascript
ServerAddress: '',
DefaultLanguage: '',
```

- [ ] **Step 2: Add `submitDefaultLanguage` function**

After `submitServerAddress` (line 315-318), add a new function:

```javascript
const submitDefaultLanguage = async () => {
await updateOptions([{ key: 'DefaultLanguage', value: inputs.DefaultLanguage || '' }]);
};
```

- [ ] **Step 3: Add language dropdown UI in the 通用设置 Card**

In the `通用设置` `Form.Section` (line 715-733), inside the `<Row>` block, after the `ServerAddress` `<Col>` closing tag (line 728) and before the `</Row>` closing tag (line 729), add a new `<Col>`:

```jsx
<Col xs={24} sm={24} md={24} lg={12} xl={12}>
<Form.Select
field='DefaultLanguage'
label={t('默认语言')}
placeholder={t('未设置时跟随浏览器语言')}
optionList={[
{ label: t('自动(跟随浏览器)'), value: '' },
{ label: '简体中文', value: 'zh-CN' },
{ label: '繁體中文', value: 'zh-TW' },
{ label: 'English', value: 'en' },
{ label: 'Français', value: 'fr' },
{ label: '日本語', value: 'ja' },
{ label: 'Русский', value: 'ru' },
{ label: 'Tiếng Việt', value: 'vi' },
]}
extraText={t(
'设置后,未登录用户和未设置语言偏好的已登录用户将强制使用此语言',
)}
/>
</Col>
```

Also change the `ServerAddress` Col from `md={24} lg={24} xl={24}` to `md={24} lg={12} xl={12}` so both fields sit side by side on large screens.

- [ ] **Step 4: Add save button for DefaultLanguage**

After the existing `submitServerAddress` button (line 730-732), add:

```jsx
<Button onClick={submitDefaultLanguage}>
{t('保存默认语言')}
</Button>
```

- [ ] **Step 5: Verify frontend compiles**

Run: `cd web && bun run build`
Expected: build succeeds

- [ ] **Step 6: Commit**

```bash
git add web/src/components/settings/SystemSetting.jsx
git commit -m "feat(default-language): add language dropdown in system settings"
```

---

### Task 4: Frontend — Apply default language for unauthenticated users

**Files:**
- Modify: `web/src/components/layout/PageLayout.jsx:87-100` (loadStatus function)

- [ ] **Step 1: Modify `loadStatus` to apply default language**

In `web/src/components/layout/PageLayout.jsx`, in the `loadStatus` function (line 87-100), after `setStatusData(data)` (line 93) and before the `} else {` (line 94), add:

```javascript
// Apply admin-configured default language for unauthenticated users
if (data.default_language && !localStorage.getItem('user')) {
i18n.changeLanguage(data.default_language);
}
```

- [ ] **Step 2: Update the existing localStorage language fallback logic**

In the same file, the existing `useEffect` (line 102-120) reads `localStorage.getItem('i18nextLng')` and calls `i18n.changeLanguage(savedLang)`. This logic should be kept as-is — it handles the case where `default_language` is not set (admin chose "auto"). No changes needed to this block.

- [ ] **Step 3: Verify frontend compiles**

Run: `cd web && bun run build`
Expected: build succeeds

- [ ] **Step 4: Commit**

```bash
git add web/src/components/layout/PageLayout.jsx
git commit -m "feat(default-language): apply default language for unauthenticated users"
```

---

### Task 5: Frontend — Fall back to default language for logged-in users without preference

**Files:**
- Modify: `web/src/context/User/index.jsx:20-45`

- [ ] **Step 1: Import StatusContext**

At the top of `web/src/context/User/index.jsx`, add `StatusContext` import after the existing imports:

```javascript
import { StatusContext } from '../Status';
```

- [ ] **Step 2: Access StatusContext inside UserProvider**

Inside `UserProvider` (line 29), before the `useEffect`, add:

```javascript
const [statusState] = React.useContext(StatusContext);
```

- [ ] **Step 3: Update language sync logic**

Replace the existing `useEffect` (lines 34-45) with:

```javascript
// Sync language preference when user data is loaded
useEffect(() => {
if (state.user?.setting) {
try {
const settings = JSON.parse(state.user.setting);
if (settings.language && settings.language !== i18n.language) {
i18n.changeLanguage(settings.language);
} else if (!settings.language && statusState.status?.default_language) {
// No personal preference — fall back to admin default
i18n.changeLanguage(statusState.status.default_language);
}
} catch (e) {
// Ignore parse errors
}
}
}, [state.user?.setting, statusState.status?.default_language, i18n]);
```

- [ ] **Step 4: Verify frontend compiles**

Run: `cd web && bun run build`
Expected: build succeeds

- [ ] **Step 5: Commit**

```bash
git add web/src/context/User/index.jsx
git commit -m "feat(default-language): fall back to default language for users without preference"
```

---

### Task 6: Integration verification

- [ ] **Step 1: Start full-stack dev server**

Terminal 1: `cd web && bun run dev`
Terminal 2: `go run main.go`

- [ ] **Step 2: Test admin setting**

1. Login as admin, navigate to Settings → System Settings (系统设置)
2. Find the "默认语言" dropdown in the 通用设置 section
3. Select "English" and click "保存默认语言"
4. Verify success toast appears
5. Refresh the page — the dropdown should still show "English"

- [ ] **Step 3: Test unauthenticated user behavior**

1. Open an incognito/private browser window
2. Visit the site
3. Expected: The page renders in English (not following browser language)
4. The header language selector still shows English as active

- [ ] **Step 4: Test logged-in user with personal preference**

1. Login as a regular user
2. Go to personal settings, set language to "Français"
3. Refresh — page should be in French (personal preference overrides admin default)

- [ ] **Step 5: Test logged-in user without personal preference**

1. Login as a user who has never set a language preference
2. Expected: Page renders in English (admin default), not browser language

- [ ] **Step 6: Test "auto" mode**

1. As admin, set "默认语言" back to "自动(跟随浏览器)"
2. Save, then test in incognito — should follow browser language again

- [ ] **Step 7: Verify /api/status returns the field**

```bash
curl -s http://localhost:3000/api/status | python -m json.tool | grep default_language
```

Expected: `"default_language": "en"` (or the set value, or `""` if auto)

+ 114
- 0
docs/superpowers/specs/2026-04-17-default-language-setting-design.md Переглянути файл

@@ -0,0 +1,114 @@
---
created: 2026-04-17
status: approved
scope: backend + frontend
files: 6
---

# 后台设置默认语言

> 管理员可在系统设置中配置全局默认语言,影响未登录用户和未设置语言偏好的已登录用户。

## 需求

- 管理员在**系统设置**页面配置一个「默认语言」
- **未登录用户**:强制使用管理员设定的默认语言,忽略浏览器语言检测
- **已登录但未设置语言偏好的用户**:使用管理员设定的默认语言
- **已登录且已设置语言偏好的用户**:使用个人设置的语言(不受默认语言影响)
- 支持的语言:zh-CN、zh-TW、en、fr、ja、ru、vi
- 可选"自动(跟随浏览器)"选项,等于未设定,保持现有行为

## 语言优先级

```
1. 用户个人设置语言(已登录 + 已设置语言偏好)
2. 管理员设定的默认语言(强制使用,忽略浏览器语言)
3. zh-CN(fallbackLng)
```

## 后端改动

### 1. `common/constants.go`

新增变量:

```go
var DefaultLanguage = "" // 空字符串=未设定,保持现有浏览器检测行为
```

### 2. `model/option.go`

在 `updateOptionMap` 的 switch 中新增 case:

```go
case "DefaultLanguage":
common.DefaultLanguage = value
```

### 3. `controller/misc.go` — `GetStatus()`

在 status 响应 JSON 中新增字段:

```go
"default_language": common.DefaultLanguage,
```

## 前端改动

### 4. `web/src/components/settings/SystemSetting.jsx`

- `inputs` state 新增 `DefaultLanguage: ''`
- 在系统设置表单中添加一个 `Select` 下拉框
- 选项:7 种语言 + 空字符串选项("自动 / 跟随浏览器")
- 与其他系统设置一起通过 `/api/option/` 保存

### 5. `web/src/components/layout/PageLayout.jsx`

在 `loadStatus` 回调中,获取 `data.default_language` 后:

```javascript
if (data.default_language && data.default_language !== '') {
const savedLang = localStorage.getItem('i18nextLng');
// 仅在没有用户个人设置语言时应用默认语言
// 注意:UserContext 会在用户加载后覆盖此设置
if (!localStorage.getItem('user')) {
i18n.changeLanguage(data.default_language);
}
}
```

同时保留现有的 `localStorage.getItem('i18nextLng')` 逻辑作为 fallback。

### 6. `web/src/context/User/index.jsx`

调整语言同步逻辑:

```javascript
// 当前:仅在有 settings.language 时切换
// 新增:无 settings.language 时,回退到 status 中的 default_language
if (settings.language) {
i18n.changeLanguage(settings.language);
} else if (statusDefaultLanguage) {
i18n.changeLanguage(statusDefaultLanguage);
}
```

需要从 StatusContext 获取 `default_language` 值。

## 涉及文件

| 文件 | 改动类型 | 改动量 |
|------|---------|-------|
| `common/constants.go` | 新增变量 | +1 行 |
| `model/option.go` | 新增 switch case | +2 行 |
| `controller/misc.go` | 新增 status 字段 | +1 行 |
| `web/src/components/settings/SystemSetting.jsx` | inputs + 表单 UI | ~30 行 |
| `web/src/components/layout/PageLayout.jsx` | loadStatus 后应用默认语言 | ~5 行 |
| `web/src/context/User/index.jsx` | 无偏好时回退到默认语言 | ~5 行 |

## 不涉及

- 不修改 i18n.js 的初始化配置
- 不修改 LanguageSelector 组件
- 不添加新的 API 端点
- 不修改数据库 schema(使用现有 options 表)

+ 236
- 0
docs/superpowers/specs/2026-04-30-log-chatid-upstreamid-design.md Переглянути файл

@@ -0,0 +1,236 @@
# Log Chat ID And Upstream ID Design

## Summary

This design adds two new usage-log fields:

- `chat_id`: extracted from the incoming request body
- `upstream_id`: extracted from the upstream response headers

The existing `logs.request_id` field keeps its current meaning and continues to represent the internal new-api request ID.

The extraction logic is centralized and writes normalized values into relay context first, then both consume logs and error logs read from the same context when persisting to `logs`.

## Problem

The current system has only one stable request identifier in usage logs: the internal `logs.request_id`.

For reconciliation work, that is not enough:

- the business request may already carry a `chat_id`
- the upstream provider may return its own request ID in response headers

Today these values are not recorded consistently in consume logs. Error logs have partial upstream request-id support, but consume logs do not have a unified path.

## Goals

- Persist request-body `chat_id` into usage logs
- Persist upstream response request ID into usage logs
- Keep success and error logs consistent
- Keep the extraction logic generic enough to support future upstream header changes
- Keep streaming support without buffering the full stream body
- Keep performance impact negligible

## Non-Goals

- No change to the meaning of existing `logs.request_id`
- No backfill of existing historical logs
- No first-version parsing of response body IDs
- No provider-specific per-channel logging branches unless the generic extractor cannot cover them

## Canonical Field Semantics

The `logs` table will have three distinct request identifiers:

- `request_id`: internal new-api request ID, already existing
- `chat_id`: business ID extracted from the incoming request body
- `upstream_id`: upstream request ID extracted from response headers

Field semantics must not overlap.

## Data Model

Add two nullable string columns to `logs`:

- `chat_id`
- `upstream_id`

Constraints:

- both fields default to empty string
- both fields use ordinary indexes
- neither field is unique

Reasoning:

- reconciliation queries need direct filtering and export
- uniqueness cannot be guaranteed across providers or retries

## Extraction Architecture

### 1. Request-side extraction

When the request body is parsed into the relay request object, the system performs a best-effort extraction of a top-level `chat_id`.

Behavior:

- extract only the business-level request `chat_id`
- do not deeply traverse nested objects
- do not fail the request if `chat_id` is missing
- store the result in relay context as `RelayInfo.ChatID`

### 2. Response-side extraction

When an upstream `http.Response` is received, before body consumption begins, the system performs a best-effort extraction of the upstream request ID and stores it in `RelayInfo.UpstreamID`.

First-version default rule chain:

1. `x-request-id`
2. `request-id`

This rule chain is intentionally generic. The extractor should be implemented as a reusable module with:

- a default candidate-header chain
- a future extension point for provider or channel-specific overrides

### 3. Persistence

Consume logs and error logs do not parse request bodies or response headers directly.

They only read:

- internal `request_id`
- `RelayInfo.ChatID`
- `RelayInfo.UpstreamID`

and persist them into `logs`.

This keeps extraction and logging decoupled.

## Data Flow

### Non-stream requests

1. request enters gateway
2. internal request ID is created as today
3. request parser extracts `chat_id`
4. relay runs normally
5. upstream response is received
6. response extractor reads `x-request-id` / `request-id`
7. final consume log or error log persists `request_id`, `chat_id`, `upstream_id`

### Stream requests

1. request parser extracts `chat_id` before relay starts
2. upstream response headers are received before stream body forwarding
3. response extractor reads `upstream_id` from headers
4. stream body is forwarded as usual
5. final consume log or error log persists `request_id`, `chat_id`, `upstream_id`

No per-chunk ID parsing is required in the first version.

## Retry Semantics

Retries must not collapse multiple upstream attempts into one combined ID set.

Rules:

- consume log records only the final successful attempt's `upstream_id`
- each failed attempt may still produce its own error log with its own `upstream_id`
- the final consume log should not store a list of retry upstream IDs

This preserves clean reconciliation semantics for billed requests.

## Performance Requirements

The design must not materially affect relay throughput or stream latency.

Allowed work:

- read response headers once
- best-effort request-side `chat_id` extraction while request parsing already happens

Forbidden first-version approaches:

- buffering the full response body only to find IDs
- parsing every SSE chunk to search for IDs
- adding extra database writes per request

Expected impact:

- request-side `chat_id` extraction: negligible, because request parsing already occurs
- response-side `upstream_id` extraction: negligible, because headers are already available
- database overhead: one existing log insert with two extra fields

## Error Handling

The feature is best-effort only.

Rules:

- missing `chat_id` does not fail the request
- missing upstream header does not fail the request
- malformed request payload does not add extra parsing failure beyond current validation behavior
- logs may store empty values for either field

## Query And UI Scope

First version should support:

- backend filtering by `chat_id`
- backend filtering by `upstream_id`
- frontend display in log detail or an equivalent low-risk display path

The first version does not need both fields as default main table columns.

## Testing Scope

### Backend tests

- non-stream success: request body includes `chat_id`, upstream returns `x-request-id`, consume log stores both
- stream success: request body includes `chat_id`, upstream stream response includes `x-request-id`, consume log stores both
- error response: error log stores both when available
- retry then success: consume log stores only the final successful `upstream_id`
- missing fields: request still succeeds or fails normally, log fields remain empty
- header fallback: if `x-request-id` is absent and `request-id` exists, `upstream_id` is still stored

### Regression focus

- no change to existing `logs.request_id` behavior
- no extra stream buffering
- no change to billing semantics

## Implementation Notes

Recommended implementation shape:

- extend `relay/common.RelayInfo` with `ChatID` and `UpstreamID`
- add a small extraction helper for request-side `chat_id`
- add a reusable response-header extractor for upstream IDs
- extend `model.Log` and log persistence functions to carry `chat_id` and `upstream_id`
- update log query/filter surfaces to support the two new fields

## Risks

- some providers may later rename or stop returning `x-request-id`
- some request formats may not include top-level `chat_id`
- stream and retry paths can drift if the extractor is not wired at a shared layer

Mitigation:

- use one shared response extractor with a candidate chain
- use one shared relay context as the single source of truth
- test both normal and retry flows explicitly

## Decision Summary

Approved design decisions:

- keep `logs.request_id` unchanged as the internal request ID
- add `logs.chat_id`
- add `logs.upstream_id`
- extract `chat_id` from the request body
- extract `upstream_id` from response headers, defaulting to `x-request-id` and falling back to `request-id`
- centralize extraction into shared relay context
- support both stream and non-stream requests without full-body buffering
- record only the final successful `upstream_id` in consume logs during retry scenarios

+ 1
- 0
dto/openai_response.go Переглянути файл

@@ -375,6 +375,7 @@ const (
type ResponsesStreamResponse struct { type ResponsesStreamResponse struct {
Type string `json:"type"` Type string `json:"type"`
Response *OpenAIResponsesResponse `json:"response,omitempty"` Response *OpenAIResponsesResponse `json:"response,omitempty"`
Error any `json:"error,omitempty"`
Delta string `json:"delta,omitempty"` Delta string `json:"delta,omitempty"`
Item *ResponsesOutput `json:"item,omitempty"` Item *ResponsesOutput `json:"item,omitempty"`
// - response.function_call_arguments.delta // - response.function_call_arguments.delta


+ 2
- 0
go.mod Переглянути файл

@@ -94,6 +94,7 @@ require (
github.com/go-sql-driver/mysql v1.7.0 // indirect github.com/go-sql-driver/mysql v1.7.0 // indirect
github.com/go-webauthn/x v0.1.25 // indirect github.com/go-webauthn/x v0.1.25 // indirect
github.com/goccy/go-json v0.10.2 // indirect github.com/goccy/go-json v0.10.2 // indirect
github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 // indirect
github.com/google/go-tpm v0.9.5 // indirect github.com/google/go-tpm v0.9.5 // indirect
github.com/gorilla/context v1.1.1 // indirect github.com/gorilla/context v1.1.1 // indirect
github.com/gorilla/securecookie v1.1.1 // indirect github.com/gorilla/securecookie v1.1.1 // indirect
@@ -117,6 +118,7 @@ require (
github.com/mitchellh/mapstructure v1.5.0 // indirect github.com/mitchellh/mapstructure v1.5.0 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/mojocn/base64Captcha v1.3.8 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/ncruces/go-strftime v0.1.9 // indirect github.com/ncruces/go-strftime v0.1.9 // indirect
github.com/pelletier/go-toml/v2 v2.2.1 // indirect github.com/pelletier/go-toml/v2 v2.2.1 // indirect


+ 57
- 0
go.sum Переглянути файл

@@ -131,11 +131,14 @@ github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I= github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo= github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 h1:DACJavvAHhabrF08vX0COfcOBJRhZ8lUbR+ZWIs0Y5g=
github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0/go.mod h1:E/TSTwGwJL78qG/PmXZO1EjYhfJinVAhrmmHX6Z8B9k=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da h1:oI5xCqsCo564l8iNU+DwB5epxmsaqB+rhGL0m5jtYqE= github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da h1:oI5xCqsCo564l8iNU+DwB5epxmsaqB+rhGL0m5jtYqE=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
github.com/golang/protobuf v1.3.3/go.mod h1:vzj43D7+SQXF/4pzW/hwtAqwc6iTitCiVSaWz5lYuqw= github.com/golang/protobuf v1.3.3/go.mod h1:vzj43D7+SQXF/4pzW/hwtAqwc6iTitCiVSaWz5lYuqw=
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk= github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/go-tpm v0.9.5 h1:ocUmnDebX54dnW+MQWGQRbdaAcJELsa6PqZhJ48KwVU= github.com/google/go-tpm v0.9.5 h1:ocUmnDebX54dnW+MQWGQRbdaAcJELsa6PqZhJ48KwVU=
@@ -223,6 +226,8 @@ github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJ
github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0= github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0=
github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
github.com/mojocn/base64Captcha v1.3.8 h1:rrN9BhCwXKS8ht1e21kvR3iTaMgf4qPC9sRoV52bqEg=
github.com/mojocn/base64Captcha v1.3.8/go.mod h1:QFZy927L8HVP3+VV5z2b1EAEiv1KxVJKZbAucVgLUy4=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
@@ -324,6 +329,7 @@ github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZ
github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E=
github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c h1:xA2TJS9Hu/ivzaZIrDcwvpJ3Fnpsk5fDOJ4iSnL6J0w= github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c h1:xA2TJS9Hu/ivzaZIrDcwvpJ3Fnpsk5fDOJ4iSnL6J0w=
github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c/go.mod h1:WSZ59bidJOO40JSJmLqlkBJrjZCtjbKKkygEMfzY/kc= github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c/go.mod h1:WSZ59bidJOO40JSJmLqlkBJrjZCtjbKKkygEMfzY/kc=
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
github.com/yusufpapurcu/wmi v1.2.3 h1:E1ctvB7uKFMOJw3fdOW32DwGE9I7t++CRUEMKvFoFiw= github.com/yusufpapurcu/wmi v1.2.3 h1:E1ctvB7uKFMOJw3fdOW32DwGE9I7t++CRUEMKvFoFiw=
github.com/yusufpapurcu/wmi v1.2.3/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0= github.com/yusufpapurcu/wmi v1.2.3/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y=
@@ -332,21 +338,46 @@ go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
golang.org/x/arch v0.21.0 h1:iTC9o7+wP6cPWpDWkivCvQFGAHDQ59SrSxsLPcnkArw= golang.org/x/arch v0.21.0 h1:iTC9o7+wP6cPWpDWkivCvQFGAHDQ59SrSxsLPcnkArw=
golang.org/x/arch v0.21.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= golang.org/x/arch v0.21.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts= golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos= golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8=
golang.org/x/image v0.23.0 h1:HseQ7c2OpPKTPVzNjG5fwJsOTCiiwS4QdsYi5XU6H68= golang.org/x/image v0.23.0 h1:HseQ7c2OpPKTPVzNjG5fwJsOTCiiwS4QdsYi5XU6H68=
golang.org/x/image v0.23.0/go.mod h1:wJJBTdLfCCf3tiHa1fNxpZmUI4mmoZvwMCPP0ddoNKY= golang.org/x/image v0.23.0/go.mod h1:wJJBTdLfCCf3tiHa1fNxpZmUI4mmoZvwMCPP0ddoNKY=
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c= golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU= golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.0.0-20210520170846-37e1c6afe023/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20210520170846-37e1c6afe023/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190726091711-fc99dfbffb4e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190726091711-fc99dfbffb4e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -355,21 +386,47 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210806184541-e5e7981a1069/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210806184541-e5e7981a1069/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
golang.org/x/term v0.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo=
golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU=
golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk=
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc= golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg= golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw= google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw=
google.golang.org/protobuf v1.28.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I= google.golang.org/protobuf v1.28.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I=


+ 5
- 3
logger/logger.go Переглянути файл

@@ -79,9 +79,11 @@ func logHelper(ctx context.Context, level string, msg string) {
if level == loggerINFO { if level == loggerINFO {
writer = gin.DefaultWriter writer = gin.DefaultWriter
} }
id := ctx.Value(common.RequestIdKey)
if id == nil {
id = "SYSTEM"
id := "SYSTEM"
if ctx != nil {
if v := ctx.Value(common.RequestIdKey); v != nil {
id = fmt.Sprintf("%v", v)
}
} }
now := time.Now() now := time.Now()
_, _ = fmt.Fprintf(writer, "[%s] %v | %s | %s \n", level, now.Format("2006/01/02 - 15:04:05"), id, msg) _, _ = fmt.Fprintf(writer, "[%s] %v | %s | %s \n", level, now.Format("2006/01/02 - 15:04:05"), id, msg)


+ 10
- 0
main.go Переглянути файл

@@ -137,6 +137,16 @@ func main() {
model.InitBatchUpdater() model.InitBatchUpdater()
} }


logoFilePath := os.Getenv("LOGO_FILE_PATH")
if logoFilePath != "" {
if _, err := os.Stat(logoFilePath); err != nil {
common.SysLog("LOGO_FILE_PATH file not found: " + logoFilePath + ", falling back to default")
} else {
common.LogoFilePath = logoFilePath
common.SysLog("custom logo file: " + common.LogoFilePath)
}
}

if os.Getenv("ENABLE_PPROF") == "true" { if os.Getenv("ENABLE_PPROF") == "true" {
gopool.Go(func() { gopool.Go(func() {
log.Println(http.ListenAndServe("0.0.0.0:8005", nil)) log.Println(http.ListenAndServe("0.0.0.0:8005", nil))


+ 62
- 14
middleware/distributor.go Переглянути файл

@@ -23,19 +23,20 @@ import (
) )


type ModelRequest struct { type ModelRequest struct {
Model string `json:"model"`
Group string `json:"group,omitempty"`
Model string `json:"model"`
Group string `json:"group,omitempty"`
ChannelId int `json:"channel_id,omitempty"`
} }


func Distribute() func(c *gin.Context) { func Distribute() func(c *gin.Context) {
return func(c *gin.Context) { return func(c *gin.Context) {
var channel *model.Channel var channel *model.Channel
channelId, ok := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId)
modelRequest, shouldSelectChannel, err := getModelRequest(c) modelRequest, shouldSelectChannel, err := getModelRequest(c)
if err != nil { if err != nil {
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()})) abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()}))
return return
} }
channelId, ok := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId)
if ok { if ok {
id, err := strconv.Atoi(channelId.(string)) id, err := strconv.Atoi(channelId.(string))
if err != nil { if err != nil {
@@ -107,25 +108,61 @@ func Distribute() func(c *gin.Context) {
} }
} }


if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found {
preferred, err := model.CacheGetChannel(preferredChannelID)
if err == nil && preferred != nil && preferred.Status == common.ChannelStatusEnabled {
if usingGroup == "auto" {
// 默认通道检查(管理员配置的模型默认通道)
if channel == nil {
if defaultChannelId, ok := model.GetDefaultChannelId(modelRequest.Model); ok {
defaultCh, err := model.CacheGetChannel(defaultChannelId)
if err != nil || defaultCh == nil {
common.SysLog(fmt.Sprintf("[Distribute] model=%s default_channel=%d not found in cache, fallback", modelRequest.Model, defaultChannelId))
} else if defaultCh.Status != common.ChannelStatusEnabled {
common.SysLog(fmt.Sprintf("[Distribute] model=%s default_channel=%d disabled(status=%d), fallback", modelRequest.Model, defaultChannelId, defaultCh.Status))
} else if usingGroup == "auto" {
userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup) userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
autoGroups := service.GetUserAutoGroup(userGroup) autoGroups := service.GetUserAutoGroup(userGroup)
for _, g := range autoGroups { for _, g := range autoGroups {
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) {
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, defaultCh.Id) {
channel = defaultCh
selectGroup = g selectGroup = g
common.SetContextKey(c, constant.ContextKeyAutoGroup, g) common.SetContextKey(c, constant.ContextKeyAutoGroup, g)
channel = preferred
service.MarkChannelAffinityUsed(c, g, preferred.Id)
common.SysLog(fmt.Sprintf("[Distribute] model=%s using default_channel=%d (auto group=%s)", modelRequest.Model, defaultChannelId, g))
break break
} }
} }
} else if model.IsChannelEnabledForGroupModel(usingGroup, modelRequest.Model, preferred.Id) {
channel = preferred
if channel == nil {
common.SysLog(fmt.Sprintf("[Distribute] model=%s default_channel=%d not enabled for any auto group, fallback", modelRequest.Model, defaultChannelId))
}
} else if model.IsChannelEnabledForGroupModel(usingGroup, modelRequest.Model, defaultCh.Id) {
channel = defaultCh
selectGroup = usingGroup selectGroup = usingGroup
service.MarkChannelAffinityUsed(c, usingGroup, preferred.Id)
common.SysLog(fmt.Sprintf("[Distribute] model=%s using default_channel=%d (group=%s)", modelRequest.Model, defaultChannelId, usingGroup))
} else {
common.SysLog(fmt.Sprintf("[Distribute] model=%s default_channel=%d not enabled for group=%s, fallback", modelRequest.Model, defaultChannelId, usingGroup))
}
}
}

// 通道亲和性检查(仅在未选中默认通道时生效)
if channel == nil {
if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found {
preferred, err := model.CacheGetChannel(preferredChannelID)
if err == nil && preferred != nil && preferred.Status == common.ChannelStatusEnabled {
if usingGroup == "auto" {
userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
autoGroups := service.GetUserAutoGroup(userGroup)
for _, g := range autoGroups {
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) {
selectGroup = g
common.SetContextKey(c, constant.ContextKeyAutoGroup, g)
channel = preferred
service.MarkChannelAffinityUsed(c, g, preferred.Id)
break
}
}
} else if model.IsChannelEnabledForGroupModel(usingGroup, modelRequest.Model, preferred.Id) {
channel = preferred
selectGroup = usingGroup
service.MarkChannelAffinityUsed(c, usingGroup, preferred.Id)
}
} }
} }
} }
@@ -330,8 +367,19 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) {
return nil, false, err return nil, false, err
} }
modelRequest.Model = req.Model modelRequest.Model = req.Model
modelRequest.Group = req.Group
// group: body 优先,fallback 到 header
if req.Group != "" {
modelRequest.Group = req.Group
} else if g := c.GetHeader("X-Group"); g != "" {
modelRequest.Group = g
}
common.SetContextKey(c, constant.ContextKeyTokenGroup, modelRequest.Group) common.SetContextKey(c, constant.ContextKeyTokenGroup, modelRequest.Group)
// channel_id: body 优先,fallback 到 header
if req.ChannelId > 0 {
common.SetContextKey(c, constant.ContextKeyTokenSpecificChannelId, strconv.Itoa(req.ChannelId))
} else if ch := c.GetHeader("X-Channel-Id"); ch != "" {
common.SetContextKey(c, constant.ContextKeyTokenSpecificChannelId, ch)
}
} }


if strings.HasPrefix(c.Request.URL.Path, "/v1/responses/compact") && modelRequest.Model != "" { if strings.HasPrefix(c.Request.URL.Path, "/v1/responses/compact") && modelRequest.Model != "" {


+ 53
- 0
model/ability.go Переглянути файл

@@ -66,6 +66,59 @@ func GetAbilitiesByChannelId(channelId int) ([]*Ability, error) {
return abilities, err return abilities, err
} }


// GetModelChannelsForGroup 返回指定模型在指定分组下的可用渠道列表及默认渠道ID
func GetModelChannelsForGroup(modelName string, group string) ([]map[string]any, int, error) {
var channelIds []int
err := DB.Model(&Ability{}).
Where("model = ?", modelName).
Where("enabled = ?", true).
Where(commonGroupCol+" = ?", group).
Distinct("channel_id").
Pluck("channel_id", &channelIds).Error
if err != nil {
return nil, 0, err
}

if len(channelIds) == 0 {
return []map[string]any{}, 0, nil
}

type channelInfo struct {
Id int `json:"id"`
Name string `json:"name"`
PublicName string `json:"public_name"`
}
var channels []channelInfo
err = DB.Table("channels").
Where("id IN ? AND status = ?", channelIds, common.ChannelStatusEnabled).
Select("id, name, public_name").
Find(&channels).Error
if err != nil {
return nil, 0, err
}

defaultChannelId := 0
if defaultChId, ok := GetDefaultChannelId(modelName); ok {
for _, id := range channelIds {
if id == defaultChId {
defaultChannelId = defaultChId
break
}
}
}

result := make([]map[string]any, 0, len(channels))
for _, ch := range channels {
result = append(result, map[string]any{
"id": ch.Id,
"name": ch.Name,
"public_name": ch.PublicName,
})
}

return result, defaultChannelId, nil
}

func getPriority(group string, model string, retry int) (int, error) { func getPriority(group string, model string, retry int) (int, error) {


var priorities []int var priorities []int


+ 9
- 1
model/channel.go Переглянути файл

@@ -26,6 +26,7 @@ type Channel struct {
TestModel *string `json:"test_model"` TestModel *string `json:"test_model"`
Status int `json:"status" gorm:"default:1"` Status int `json:"status" gorm:"default:1"`
Name string `json:"name" gorm:"index"` Name string `json:"name" gorm:"index"`
PublicName string `json:"public_name" gorm:"size:255;default:''"`
Weight *uint `json:"weight" gorm:"default:0"` Weight *uint `json:"weight" gorm:"default:0"`
CreatedTime int64 `json:"created_time" gorm:"bigint"` CreatedTime int64 `json:"created_time" gorm:"bigint"`
TestTime int64 `json:"test_time" gorm:"bigint"` TestTime int64 `json:"test_time" gorm:"bigint"`
@@ -72,6 +73,13 @@ func (c ChannelInfo) Value() (driver.Value, error) {
return common.Marshal(&c) return common.Marshal(&c)
} }


func ChannelDisplayName(publicName, name string) string {
if publicName != "" {
return publicName
}
return name
}

// Scan implements sql.Scanner interface // Scan implements sql.Scanner interface
func (c *ChannelInfo) Scan(value interface{}) error { func (c *ChannelInfo) Scan(value interface{}) error {
bytesValue, _ := value.([]byte) bytesValue, _ := value.([]byte)
@@ -279,7 +287,7 @@ func GetAllChannels(startIdx int, num int, selectAll bool, idSort bool) ([]*Chan
// 只返回 id, name, type, remark,不包含敏感信息 // 只返回 id, name, type, remark,不包含敏感信息
func GetAllChannelsForBinding() ([]*Channel, error) { func GetAllChannelsForBinding() ([]*Channel, error) {
var channels []*Channel var channels []*Channel
err := DB.Select("id, name, type, remark").
err := DB.Select("id, name, public_name, type, remark").
Where("status = ?", common.ChannelStatusEnabled). Where("status = ?", common.ChannelStatusEnabled).
Order("priority desc"). Order("priority desc").
Find(&channels).Error Find(&channels).Error


+ 226
- 79
model/channel_pricing.go Переглянути файл

@@ -5,7 +5,6 @@ import (
"strconv" "strconv"
"strings" "strings"
"sync" "sync"
"time"


"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/QuantumNous/new-api/setting/ratio_setting"
@@ -17,8 +16,10 @@ import (
var ( var (
channelPricingCache = make(map[string]*ChannelPricing) // key: "modelName:channelId" channelPricingCache = make(map[string]*ChannelPricing) // key: "modelName:channelId"
channelPricingCacheLock sync.RWMutex channelPricingCacheLock sync.RWMutex
channelPricingCacheTime time.Time
channelPricingCacheTTL = time.Minute * 5 // 缓存5分钟

// 默认通道缓存:modelName → channelId
defaultChannelCache = make(map[string]int)
defaultChannelCacheLock sync.RWMutex
) )


// QuotaType 计费类型 // QuotaType 计费类型
@@ -41,6 +42,40 @@ type ChannelPricing struct {
CreatedTime int64 `json:"created_time" gorm:"bigint"` CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"` UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
DeletedAt gorm.DeletedAt `json:"-" gorm:"index"` DeletedAt gorm.DeletedAt `json:"-" gorm:"index"`
// === 新增字段(0 = 未设置,回退全局值) ===
CacheRatio float64 `json:"cache_ratio" gorm:"default:0"`
CacheCreationRatio float64 `json:"cache_creation_ratio" gorm:"default:0"`
ImageRatio float64 `json:"image_ratio" gorm:"default:0"`
AudioRatio float64 `json:"audio_ratio" gorm:"default:0"`
AudioCompletionRatio float64 `json:"audio_completion_ratio" gorm:"default:0"`
IsDefault bool `json:"is_default" gorm:"default:false;index"`
}

// setCache 写穿透缓存
func setCache(key string, cp *ChannelPricing) {
channelPricingCacheLock.Lock()
channelPricingCache[key] = cp
channelPricingCacheLock.Unlock()
}

func removeCache(key string) {
channelPricingCacheLock.Lock()
delete(channelPricingCache, key)
channelPricingCacheLock.Unlock()
}

// ApplyFields 批量设置定价字段(消除 controller 层的重复赋值)
func (cp *ChannelPricing) ApplyFields(quotaType int, modelRatio, completionRatio, modelPrice float64, tagIds string, cacheRatio, cacheCreationRatio, imageRatio, audioRatio, audioCompletionRatio float64) {
cp.QuotaType = quotaType
cp.ModelRatio = modelRatio
cp.CompletionRatio = completionRatio
cp.ModelPrice = modelPrice
cp.TagIds = tagIds
cp.CacheRatio = cacheRatio
cp.CacheCreationRatio = cacheCreationRatio
cp.ImageRatio = imageRatio
cp.AudioRatio = audioRatio
cp.AudioCompletionRatio = audioCompletionRatio
} }


func (cp *ChannelPricing) Insert() error { func (cp *ChannelPricing) Insert() error {
@@ -49,31 +84,43 @@ func (cp *ChannelPricing) Insert() error {
cp.UpdatedTime = now cp.UpdatedTime = now
err := DB.Create(cp).Error err := DB.Create(cp).Error
if err == nil { if err == nil {
InvalidateChannelPricingCache()
setCache(getChannelPricingCacheKey(cp.ModelName, cp.ChannelId), cp)
if cp.IsDefault {
setDefaultChannelCache(cp.ModelName, cp.ChannelId)
}
} }
return err return err
} }


func (cp *ChannelPricing) Update() error { func (cp *ChannelPricing) Update() error {
cp.UpdatedTime = common.GetTimestamp() cp.UpdatedTime = common.GetTimestamp()
err := DB.Model(&ChannelPricing{}).Where("id = ?", cp.Id).Updates(map[string]interface{}{
"quota_type": cp.QuotaType,
"model_ratio": cp.ModelRatio,
"completion_ratio": cp.CompletionRatio,
"model_price": cp.ModelPrice,
"tag_ids": cp.TagIds,
"updated_time": cp.UpdatedTime,
}).Error
err := DB.Model(&ChannelPricing{}).Where("id = ?", cp.Id).
Select("quota_type", "model_ratio", "completion_ratio", "model_price",
"tag_ids", "cache_ratio", "cache_creation_ratio", "image_ratio",
"audio_ratio", "audio_completion_ratio", "is_default", "updated_time").
Updates(cp).Error
if err == nil { if err == nil {
InvalidateChannelPricingCache()
setCache(getChannelPricingCacheKey(cp.ModelName, cp.ChannelId), cp)
if cp.IsDefault {
setDefaultChannelCache(cp.ModelName, cp.ChannelId)
} else {
clearDefaultChannelCacheIfMatch(cp.ModelName, cp.Id)
}
} }
return err return err
} }


func (cp *ChannelPricing) Delete() error { func (cp *ChannelPricing) Delete() error {
var existing ChannelPricing
if err := DB.First(&existing, cp.Id).Error; err != nil {
return err
}
err := DB.Delete(cp).Error err := DB.Delete(cp).Error
if err == nil { if err == nil {
InvalidateChannelPricingCache()
removeCache(getChannelPricingCacheKey(existing.ModelName, existing.ChannelId))
if existing.IsDefault {
clearDefaultChannelCache(existing.ModelName)
}
} }
return err return err
} }
@@ -121,7 +168,7 @@ func BatchUpsertChannelPricing(pricings []*ChannelPricing) error {
} }
// 使用 GORM 的 OnConflict 实现 upsert // 使用 GORM 的 OnConflict 实现 upsert
// 唯一索引为 idx_model_channel (model_name, channel_id) // 唯一索引为 idx_model_channel (model_name, channel_id)
return DB.Clauses(clause.OnConflict{
err := DB.Clauses(clause.OnConflict{
Columns: []clause.Column{ Columns: []clause.Column{
{Name: "model_name"}, {Name: "model_name"},
{Name: "channel_id"}, {Name: "channel_id"},
@@ -132,9 +179,20 @@ func BatchUpsertChannelPricing(pricings []*ChannelPricing) error {
"completion_ratio", "completion_ratio",
"model_price", "model_price",
"tag_ids", "tag_ids",
"cache_ratio",
"cache_creation_ratio",
"image_ratio",
"audio_ratio",
"audio_completion_ratio",
"updated_time", "updated_time",
}), }),
}).Create(&pricings).Error }).Create(&pricings).Error
if err == nil {
for _, cp := range pricings {
setCache(getChannelPricingCacheKey(cp.ModelName, cp.ChannelId), cp)
}
}
return err
} }


// getChannelPricingCacheKey 生成缓存键 // getChannelPricingCacheKey 生成缓存键
@@ -142,71 +200,51 @@ func getChannelPricingCacheKey(modelName string, channelId int) string {
return fmt.Sprintf("%s:%d", modelName, channelId) return fmt.Sprintf("%s:%d", modelName, channelId)
} }


// GetEffectivePricing 获取有效定价(优先渠道定价,回退全局定价)
// 返回: modelRatio, completionRatio, modelPrice, usePrice, found
func GetEffectivePricing(modelName string, channelId int) (modelRatio, completionRatio, modelPrice float64, usePrice, found bool) {
cacheKey := getChannelPricingCacheKey(modelName, channelId)

// 首先检查缓存
// GetEffectivePricing 获取有效定价(纯内存查找)
func GetEffectivePricing(modelName string, channelId int) (*ChannelPricing, bool) {
key := getChannelPricingCacheKey(modelName, channelId)
channelPricingCacheLock.RLock() channelPricingCacheLock.RLock()
// 检查缓存是否过期
if time.Since(channelPricingCacheTime) < channelPricingCacheTTL {
if cp, ok := channelPricingCache[cacheKey]; ok {
channelPricingCacheLock.RUnlock()
return cp.ModelRatio, cp.CompletionRatio, cp.ModelPrice, cp.QuotaType == QuotaTypeByCall, true
}
}
cp, ok := channelPricingCache[key]
channelPricingCacheLock.RUnlock() channelPricingCacheLock.RUnlock()

// 缓存未命中或已过期,查询数据库
var cp ChannelPricing
err := DB.Where("model_name = ? AND channel_id = ?", modelName, channelId).First(&cp).Error
if err != nil {
// 未找到渠道定价,返回 false 让调用者使用全局定价
return 0, 0, 0, false, false
if !ok {
return nil, false
} }
return cp, true
}


// 更新缓存
channelPricingCacheLock.Lock()
if channelPricingCacheTime.IsZero() || time.Since(channelPricingCacheTime) >= channelPricingCacheTTL {
// 缓存过期,清空并更新时间
channelPricingCache = make(map[string]*ChannelPricing)
channelPricingCacheTime = time.Now()
// ParseTagIds 解析逗号分隔的标签ID字符串为 PricingTag 切片
func ParseTagIds(tagIds string, tagMap map[int]*PricingTag) []*PricingTag {
if tagIds == "" {
return nil
} }
channelPricingCache[cacheKey] = &cp
channelPricingCacheLock.Unlock()

return cp.ModelRatio, cp.CompletionRatio, cp.ModelPrice, cp.QuotaType == QuotaTypeByCall, true
tags := make([]*PricingTag, 0)
for _, idStr := range strings.Split(tagIds, ",") {
if id, err := strconv.Atoi(strings.TrimSpace(idStr)); err == nil {
if tag, ok := tagMap[id]; ok {
tags = append(tags, tag)
}
}
}
return tags
} }


// RefreshChannelPricingCache 刷新渠道定价缓存
func RefreshChannelPricingCache() {
channelPricingCacheLock.Lock()
defer channelPricingCacheLock.Unlock()

// 清空缓存
channelPricingCache = make(map[string]*ChannelPricing)
channelPricingCacheTime = time.Now()

// 预加载所有渠道定价
// LoadChannelPricingCache 全量加载渠道定价到内存(启动时调用)
func LoadChannelPricingCache() {
var pricings []*ChannelPricing var pricings []*ChannelPricing
if err := DB.Find(&pricings).Error; err != nil { if err := DB.Find(&pricings).Error; err != nil {
common.SysError("[ChannelPricing] LoadChannelPricingCache failed: " + err.Error())
return return
} }

channelPricingCacheLock.Lock()
channelPricingCache = make(map[string]*ChannelPricing, len(pricings))
for _, cp := range pricings { for _, cp := range pricings {
cacheKey := getChannelPricingCacheKey(cp.ModelName, cp.ChannelId)
channelPricingCache[cacheKey] = cp
key := getChannelPricingCacheKey(cp.ModelName, cp.ChannelId)
channelPricingCache[key] = cp
} }
}

// InvalidateChannelPricingCache 使渠道定价缓存失效
func InvalidateChannelPricingCache() {
channelPricingCacheLock.Lock()
defer channelPricingCacheLock.Unlock()
channelPricingCacheLock.Unlock()


channelPricingCache = make(map[string]*ChannelPricing)
channelPricingCacheTime = time.Time{} // 重置为零值
rebuildDefaultChannelCache(pricings)
common.SysLog(fmt.Sprintf("[ChannelPricing] cache loaded %d records", len(pricings)))
} }


// ChannelPricingWithChannel 带渠道信息的定价响应 // ChannelPricingWithChannel 带渠道信息的定价响应
@@ -214,6 +252,7 @@ type ChannelPricingWithChannel struct {
Id int `json:"id"` Id int `json:"id"`
ChannelId int `json:"channel_id"` ChannelId int `json:"channel_id"`
ChannelName string `json:"channel_name"` ChannelName string `json:"channel_name"`
ChannelPublicName string `json:"channel_public_name"`
ChannelType int `json:"channel_type"` ChannelType int `json:"channel_type"`
TagIds string `json:"tag_ids" gorm:"column:tag_ids"` // 渠道定价的标签ID列表(逗号分隔) TagIds string `json:"tag_ids" gorm:"column:tag_ids"` // 渠道定价的标签ID列表(逗号分隔)
Tags []*PricingTag `json:"tags" gorm:"-"` // 渠道定价的标签详情(不参与数据库扫描) Tags []*PricingTag `json:"tags" gorm:"-"` // 渠道定价的标签详情(不参与数据库扫描)
@@ -221,7 +260,13 @@ type ChannelPricingWithChannel struct {
ModelRatio float64 `json:"model_ratio"` ModelRatio float64 `json:"model_ratio"`
CompletionRatio float64 `json:"completion_ratio"` CompletionRatio float64 `json:"completion_ratio"`
ModelPrice float64 `json:"model_price"` ModelPrice float64 `json:"model_price"`
HasCustomPricing bool `json:"has_custom_pricing"` // 是否有自定义定价
HasCustomPricing bool `json:"has_custom_pricing"` // 是否有自定义定价
CacheRatio float64 `json:"cache_ratio"`
CacheCreationRatio float64 `json:"cache_creation_ratio"`
ImageRatio float64 `json:"image_ratio"`
AudioRatio float64 `json:"audio_ratio"`
AudioCompletionRatio float64 `json:"audio_completion_ratio"`
IsDefault bool `json:"is_default"`
} }


// GetChannelPricingByModelWithChannelInfo 获取指定模型的渠道定价(带渠道信息) // GetChannelPricingByModelWithChannelInfo 获取指定模型的渠道定价(带渠道信息)
@@ -249,16 +294,22 @@ func GetChannelPricingByModelWithChannelInfo(modelName string) ([]*ChannelPricin
if !hasPrice { if !hasPrice {
globalModelPrice = 0 globalModelPrice = 0
} }

// 查询所有支持该模型的渠道,左连接渠道定价表 // 查询所有支持该模型的渠道,左连接渠道定价表
// 高级字段(cache/image/audio)不回退全局值,直接返回 0
err := DB.Table("abilities"). err := DB.Table("abilities").
Select(`abilities.channel_id, channels.name as channel_name, channels.type as channel_type,
Select(`abilities.channel_id, channels.name as channel_name, channels.public_name as channel_public_name, channels.type as channel_type,
COALESCE(channel_pricings.quota_type, ?) as quota_type, COALESCE(channel_pricings.quota_type, ?) as quota_type,
COALESCE(channel_pricings.model_ratio, ?) as model_ratio, COALESCE(channel_pricings.model_ratio, ?) as model_ratio,
COALESCE(channel_pricings.completion_ratio, ?) as completion_ratio, COALESCE(channel_pricings.completion_ratio, ?) as completion_ratio,
COALESCE(channel_pricings.model_price, ?) as model_price, COALESCE(channel_pricings.model_price, ?) as model_price,
channel_pricings.id as id, channel_pricings.id as id,
channel_pricings.tag_ids as tag_ids, channel_pricings.tag_ids as tag_ids,
COALESCE(channel_pricings.cache_ratio, 0) as cache_ratio,
COALESCE(channel_pricings.cache_creation_ratio, 0) as cache_creation_ratio,
COALESCE(channel_pricings.image_ratio, 0) as image_ratio,
COALESCE(channel_pricings.audio_ratio, 0) as audio_ratio,
COALESCE(channel_pricings.audio_completion_ratio, 0) as audio_completion_ratio,
COALESCE(channel_pricings.is_default, false) as is_default,
(channel_pricings.id IS NOT NULL) as has_custom_pricing`, (channel_pricings.id IS NOT NULL) as has_custom_pricing`,
defaultQuotaType, globalModelRatio, globalCompletionRatio, globalModelPrice). defaultQuotaType, globalModelRatio, globalCompletionRatio, globalModelPrice).
Joins("LEFT JOIN channels ON abilities.channel_id = channels.id"). Joins("LEFT JOIN channels ON abilities.channel_id = channels.id").
@@ -266,7 +317,7 @@ func GetChannelPricingByModelWithChannelInfo(modelName string) ([]*ChannelPricin
Where("abilities.model = ?", modelName). Where("abilities.model = ?", modelName).
Where("abilities.enabled = ?", true). Where("abilities.enabled = ?", true).
Where("channels.status = ?", 1). // 只显示启用的渠道 Where("channels.status = ?", 1). // 只显示启用的渠道
Group("abilities.channel_id, channels.name, channels.type, channel_pricings.quota_type, channel_pricings.model_ratio, channel_pricings.completion_ratio, channel_pricings.model_price, channel_pricings.id, channel_pricings.tag_ids").
Group("abilities.channel_id, channels.name, channels.public_name, channels.type, channel_pricings.quota_type, channel_pricings.model_ratio, channel_pricings.completion_ratio, channel_pricings.model_price, channel_pricings.id, channel_pricings.tag_ids, channel_pricings.cache_ratio, channel_pricings.cache_creation_ratio, channel_pricings.image_ratio, channel_pricings.audio_ratio, channel_pricings.audio_completion_ratio, channel_pricings.is_default").
Scan(&results).Error Scan(&results).Error
if err != nil { if err != nil {
return nil, err return nil, err
@@ -286,17 +337,113 @@ func GetChannelPricingByModelWithChannelInfo(modelName string) ([]*ChannelPricin


// 为每个渠道定价填充标签 // 为每个渠道定价填充标签
for _, result := range results { for _, result := range results {
if result.TagIds != "" {
result.Tags = make([]*PricingTag, 0)
for _, idStr := range strings.Split(result.TagIds, ",") {
if id, err := strconv.Atoi(strings.TrimSpace(idStr)); err == nil {
if tag, ok := tagMap[id]; ok {
result.Tags = append(result.Tags, tag)
}
}
result.Tags = ParseTagIds(result.TagIds, tagMap)
}

return results, nil
}

// === 默认通道缓存 ===

// syncIsDefaultToCache 将默认标记的变更同步到 channelPricingCache
func syncIsDefaultToCache(modelName string, channelId int, isDefault bool) {
channelPricingCacheLock.Lock()
for key, cp := range channelPricingCache {
if cp.ModelName == modelName {
if isDefault {
cp.IsDefault = cp.ChannelId == channelId
} else {
cp.IsDefault = false
} }
} }
channelPricingCache[key] = cp
} }
channelPricingCacheLock.Unlock()
}


return results, nil
// setDefaultChannelCache 设置默认通道缓存
func setDefaultChannelCache(modelName string, channelId int) {
defaultChannelCacheLock.Lock()
defaultChannelCache[modelName] = channelId
defaultChannelCacheLock.Unlock()
}

// clearDefaultChannelCache 清除指定模型的默认通道缓存
func clearDefaultChannelCache(modelName string) {
defaultChannelCacheLock.Lock()
delete(defaultChannelCache, modelName)
defaultChannelCacheLock.Unlock()
}

// clearDefaultChannelCacheIfMatch 如果默认通道的定价记录 ID 匹配则清除
func clearDefaultChannelCacheIfMatch(modelName string, pricingId int) {
defaultChannelCacheLock.RLock()
cachedId, ok := defaultChannelCache[modelName]
defaultChannelCacheLock.RUnlock()
if !ok {
return
}
// 需要通过缓存找到对应的 pricing 来比对
key := getChannelPricingCacheKey(modelName, cachedId)
channelPricingCacheLock.RLock()
cp, exists := channelPricingCache[key]
channelPricingCacheLock.RUnlock()
if exists && cp.Id == pricingId {
clearDefaultChannelCache(modelName)
}
}

// rebuildDefaultChannelCache 从全量数据构建默认通道缓存(启动时调用)
func rebuildDefaultChannelCache(pricings []*ChannelPricing) {
defaultChannelCacheLock.Lock()
defaultChannelCache = make(map[string]int)
for _, cp := range pricings {
if cp.IsDefault {
defaultChannelCache[cp.ModelName] = cp.ChannelId
}
}
defaultChannelCacheLock.Unlock()
common.SysLog(fmt.Sprintf("[ChannelPricing] default channel cache loaded %d records", len(defaultChannelCache)))
}

// GetDefaultChannelId 获取指定模型的默认通道 ID(纯内存读)
func GetDefaultChannelId(modelName string) (int, bool) {
defaultChannelCacheLock.RLock()
id, ok := defaultChannelCache[modelName]
defaultChannelCacheLock.RUnlock()
return id, ok
}

// SetDefaultChannel 设置指定模型的默认通道(事务保证互斥)
func SetDefaultChannel(modelName string, channelId int) error {
return DB.Transaction(func(tx *gorm.DB) error {
// 清除该模型所有现有的默认标记
if err := tx.Model(&ChannelPricing{}).
Where("model_name = ? AND is_default = ?", modelName, true).
Update("is_default", false).Error; err != nil {
return err
}
// 设置新的默认
if err := tx.Model(&ChannelPricing{}).
Where("model_name = ? AND channel_id = ?", modelName, channelId).
Update("is_default", true).Error; err != nil {
return err
}
// 更新缓存
setDefaultChannelCache(modelName, channelId)
syncIsDefaultToCache(modelName, channelId, true)
return nil
})
}

// ClearDefaultChannel 清除指定模型的默认通道标记
func ClearDefaultChannel(modelName string) error {
err := DB.Model(&ChannelPricing{}).
Where("model_name = ? AND is_default = ?", modelName, true).
Update("is_default", false).Error
if err == nil {
clearDefaultChannelCache(modelName)
syncIsDefaultToCache(modelName, 0, false)
}
return err
} }

+ 242
- 0
model/channel_pricing_test.go Переглянути файл

@@ -0,0 +1,242 @@
package model

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func setupChannelPricingDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := DB
DB = db
common.UsingSQLite = true
common.RedisEnabled = false

require.NoError(t, db.AutoMigrate(&ChannelPricing{}))

t.Cleanup(func() {
DB = origDB
sqlDB.Close()
})
return db
}

func TestCacheWriteThrough_Insert(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{
ModelName: "test-insert-model",
ChannelId: 9001,
QuotaType: QuotaTypeByTokens,
ModelRatio: 1.0,
CacheRatio: 0.8,
ImageRatio: 1.2,
}
require.NoError(t, cp.Insert())
defer cp.Delete()

found, ok := GetEffectivePricing("test-insert-model", 9001)
assert.True(t, ok)
assert.Equal(t, 0.8, found.CacheRatio)
assert.Equal(t, 1.2, found.ImageRatio)
}

func TestCacheWriteThrough_Update(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{
ModelName: "test-update-model",
ChannelId: 9002,
QuotaType: QuotaTypeByTokens,
ModelRatio: 1.0,
}
require.NoError(t, cp.Insert())
defer cp.Delete()

cp.CacheRatio = 0.9
cp.AudioRatio = 1.5
require.NoError(t, cp.Update())

found, ok := GetEffectivePricing("test-update-model", 9002)
assert.True(t, ok)
assert.Equal(t, 0.9, found.CacheRatio)
assert.Equal(t, 1.5, found.AudioRatio)
}

func TestCacheWriteThrough_Delete(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{
ModelName: "test-delete-model",
ChannelId: 9003,
QuotaType: QuotaTypeByTokens,
ModelRatio: 1.0,
}
require.NoError(t, cp.Insert())

require.NoError(t, cp.Delete())

found, ok := GetEffectivePricing("test-delete-model", 9003)
assert.False(t, ok)
assert.Nil(t, found)
}

func TestExtendedFields_DefaultZero(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{
ModelName: "test-default-model",
ChannelId: 9004,
QuotaType: QuotaTypeByTokens,
ModelRatio: 1.0,
}
require.NoError(t, cp.Insert())
defer cp.Delete()

found, ok := GetEffectivePricing("test-default-model", 9004)
assert.True(t, ok)
assert.Equal(t, 0.0, found.CacheRatio)
assert.Equal(t, 0.0, found.CacheCreationRatio)
assert.Equal(t, 0.0, found.ImageRatio)
assert.Equal(t, 0.0, found.AudioRatio)
assert.Equal(t, 0.0, found.AudioCompletionRatio)
}

func TestGetEffectivePricing_NotFound(t *testing.T) {
setupChannelPricingDB(t)

found, ok := GetEffectivePricing("nonexistent-model-xyz", 99999)
assert.False(t, ok)
assert.Nil(t, found)
}

// === 默认通道测试 ===

func TestDefaultChannel_SetAndGet(t *testing.T) {
setupChannelPricingDB(t)

// 创建两条定价记录
cp1 := &ChannelPricing{ModelName: "default-test-model", ChannelId: 100, QuotaType: QuotaTypeByTokens, ModelRatio: 1.0}
cp2 := &ChannelPricing{ModelName: "default-test-model", ChannelId: 200, QuotaType: QuotaTypeByTokens, ModelRatio: 2.0}
require.NoError(t, cp1.Insert())
require.NoError(t, cp2.Insert())
t.Cleanup(func() { cp1.Delete(); cp2.Delete() })

// 初始没有默认
_, ok := GetDefaultChannelId("default-test-model")
assert.False(t, ok)

// 设置通道 100 为默认
require.NoError(t, SetDefaultChannel("default-test-model", 100))
id, ok := GetDefaultChannelId("default-test-model")
assert.True(t, ok)
assert.Equal(t, 100, id)

// 切换默认到通道 200
require.NoError(t, SetDefaultChannel("default-test-model", 200))
id, ok = GetDefaultChannelId("default-test-model")
assert.True(t, ok)
assert.Equal(t, 200, id)

// 验证旧的默认标记被清除
found, _ := GetEffectivePricing("default-test-model", 100)
assert.False(t, found.IsDefault)
found, _ = GetEffectivePricing("default-test-model", 200)
assert.True(t, found.IsDefault)
}

func TestDefaultChannel_Clear(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{ModelName: "clear-test-model", ChannelId: 300, QuotaType: QuotaTypeByTokens, ModelRatio: 1.0}
require.NoError(t, cp.Insert())
t.Cleanup(func() { cp.Delete() })

require.NoError(t, SetDefaultChannel("clear-test-model", 300))
_, ok := GetDefaultChannelId("clear-test-model")
assert.True(t, ok)

require.NoError(t, ClearDefaultChannel("clear-test-model"))
_, ok = GetDefaultChannelId("clear-test-model")
assert.False(t, ok)
}

func TestDefaultChannel_DeleteClearsCache(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{ModelName: "delete-test-model", ChannelId: 400, QuotaType: QuotaTypeByTokens, ModelRatio: 1.0, IsDefault: true}
require.NoError(t, cp.Insert())

id, ok := GetDefaultChannelId("delete-test-model")
assert.True(t, ok)
assert.Equal(t, 400, id)

// 删除记录后缓存应清除
require.NoError(t, cp.Delete())
_, ok = GetDefaultChannelId("delete-test-model")
assert.False(t, ok)
}

func TestDefaultChannel_LoadCache(t *testing.T) {
db := setupChannelPricingDB(t)

// 直接插入数据(绕过缓存)
db.Create(&ChannelPricing{ModelName: "load-model", ChannelId: 500, ModelRatio: 1.0, IsDefault: true})
db.Create(&ChannelPricing{ModelName: "load-model", ChannelId: 501, ModelRatio: 2.0, IsDefault: false})
db.Create(&ChannelPricing{ModelName: "other-model", ChannelId: 502, ModelRatio: 1.0, IsDefault: true})

// 全量加载缓存
LoadChannelPricingCache()

id, ok := GetDefaultChannelId("load-model")
assert.True(t, ok)
assert.Equal(t, 500, id)

id, ok = GetDefaultChannelId("other-model")
assert.True(t, ok)
assert.Equal(t, 502, id)

// 无默认的模型
_, ok = GetDefaultChannelId("nonexistent")
assert.False(t, ok)
}

func TestDefaultChannel_InsertWithDefault(t *testing.T) {
setupChannelPricingDB(t)

cp := &ChannelPricing{ModelName: "insert-default-model", ChannelId: 600, ModelRatio: 1.0, IsDefault: true}
require.NoError(t, cp.Insert())
t.Cleanup(func() { cp.Delete() })

id, ok := GetDefaultChannelId("insert-default-model")
assert.True(t, ok)
assert.Equal(t, 600, id)
}

func TestDefaultChannel_UpdateWithDefault(t *testing.T) {
setupChannelPricingDB(t)

cp1 := &ChannelPricing{ModelName: "update-default-model", ChannelId: 700, ModelRatio: 1.0, IsDefault: true}
cp2 := &ChannelPricing{ModelName: "update-default-model", ChannelId: 701, ModelRatio: 2.0}
require.NoError(t, cp1.Insert())
require.NoError(t, cp2.Insert())
t.Cleanup(func() { cp1.Delete(); cp2.Delete() })

// cp1 是默认,通过 Update 把 cp2 设为默认
cp2.IsDefault = true
require.NoError(t, cp2.Update())

id, ok := GetDefaultChannelId("update-default-model")
assert.True(t, ok)
assert.Equal(t, 701, id)
}

+ 133
- 0
model/email_quota_rule.go Переглянути файл

@@ -0,0 +1,133 @@
package model

import (
"strings"
"sync"

"github.com/QuantumNous/new-api/common"
)

type EmailQuotaRule struct {
Id int `json:"id" gorm:"primaryKey"`
EmailSuffix string `json:"email_suffix" gorm:"size:128;not null;uniqueIndex"`
Quota int64 `json:"quota" gorm:"not null"`
Enabled bool `json:"enabled" gorm:"default:1"`
Description string `json:"description" gorm:"size:256"`
CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
}

var (
emailQuotaCache map[string]int64
emailQuotaCacheMu sync.RWMutex
)

// LoadEmailQuotaCache loads all enabled rules into memory cache.
func LoadEmailQuotaCache() {
var rules []EmailQuotaRule
DB.Where("enabled = ?", true).Find(&rules)
cache := make(map[string]int64, len(rules))
for _, r := range rules {
cache[strings.ToLower(r.EmailSuffix)] = r.Quota
}
emailQuotaCacheMu.Lock()
emailQuotaCache = cache
emailQuotaCacheMu.Unlock()
}

// MatchEmailQuotaRule checks if the email matches any enabled suffix rule.
// Returns the matched quota, or -1 if no match.
func MatchEmailQuotaRule(email string) int64 {
if email == "" {
return -1
}
at := strings.LastIndex(email, "@")
if at < 0 {
return -1
}
suffix := strings.ToLower(email[at:])
emailQuotaCacheMu.RLock()
defer emailQuotaCacheMu.RUnlock()
if quota, ok := emailQuotaCache[suffix]; ok {
return quota
}
return -1
}

func boolToInt(b bool) int {
if b {
return 1
}
return 0
}

func (r *EmailQuotaRule) Insert() error {
r.CreatedTime = common.GetTimestamp()
r.UpdatedTime = r.CreatedTime
err := DB.Model(&EmailQuotaRule{}).Create(map[string]interface{}{
"email_suffix": r.EmailSuffix,
"quota": r.Quota,
"enabled": boolToInt(r.Enabled),
"description": r.Description,
"created_time": r.CreatedTime,
"updated_time": r.UpdatedTime,
}).Error
if err != nil {
return err
}
var last EmailQuotaRule
if err := DB.Where("email_suffix = ?", r.EmailSuffix).First(&last).Error; err == nil {
r.Id = last.Id
}
LoadEmailQuotaCache()
return nil
}

func (r *EmailQuotaRule) Update() error {
r.UpdatedTime = common.GetTimestamp()
err := DB.Model(&EmailQuotaRule{}).Where("id = ?", r.Id).Updates(map[string]interface{}{
"email_suffix": r.EmailSuffix,
"quota": r.Quota,
"enabled": boolToInt(r.Enabled),
"description": r.Description,
"updated_time": r.UpdatedTime,
}).Error
if err != nil {
return err
}
LoadEmailQuotaCache()
return nil
}

func (r *EmailQuotaRule) Delete() error {
err := DB.Delete(r).Error
if err != nil {
return err
}
LoadEmailQuotaCache()
return nil
}

func GetAllEmailQuotaRules() ([]EmailQuotaRule, error) {
var list []EmailQuotaRule
err := DB.Order("id ASC").Find(&list).Error
return list, err
}

func GetEmailQuotaRuleById(id int) (*EmailQuotaRule, error) {
var r EmailQuotaRule
err := DB.First(&r, id).Error
if err != nil {
return nil, err
}
return &r, nil
}

func GetEmailQuotaRuleBySuffix(suffix string) (*EmailQuotaRule, error) {
var r EmailQuotaRule
err := DB.Where("email_suffix = ?", suffix).First(&r).Error
if err != nil {
return nil, err
}
return &r, nil
}

+ 246
- 0
model/email_quota_rule_test.go Переглянути файл

@@ -0,0 +1,246 @@
package model

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func setupEmailQuotaRuleDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := DB
DB = db
common.UsingSQLite = true
common.RedisEnabled = false

require.NoError(t, db.AutoMigrate(&EmailQuotaRule{}))

t.Cleanup(func() {
DB = origDB
sqlDB.Close()
})
return db
}

func TestMatchEmailQuotaRule_EmptyEmail(t *testing.T) {
setupEmailQuotaRuleDB(t)
LoadEmailQuotaCache()

result := MatchEmailQuotaRule("")
assert.Equal(t, int64(-1), result)
}

func TestMatchEmailQuotaRule_NoAtSign(t *testing.T) {
setupEmailQuotaRuleDB(t)
LoadEmailQuotaCache()

result := MatchEmailQuotaRule("invalidemail")
assert.Equal(t, int64(-1), result)
}

func TestMatchEmailQuotaRule_NoRules(t *testing.T) {
setupEmailQuotaRuleDB(t)
LoadEmailQuotaCache()

result := MatchEmailQuotaRule("user@example.com")
assert.Equal(t, int64(-1), result)
}

func TestMatchEmailQuotaRule_MatchEnabled(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{
EmailSuffix: "@example.com",
Quota: 500000,
Enabled: true,
}
require.NoError(t, rule.Insert())

result := MatchEmailQuotaRule("user@example.com")
assert.Equal(t, int64(500000), result)
}

func TestMatchEmailQuotaRule_CaseInsensitive(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{
EmailSuffix: "@Example.COM",
Quota: 300000,
Enabled: true,
}
require.NoError(t, rule.Insert())

assert.Equal(t, int64(300000), MatchEmailQuotaRule("user@example.com"))
assert.Equal(t, int64(300000), MatchEmailQuotaRule("user@EXAMPLE.COM"))
assert.Equal(t, int64(300000), MatchEmailQuotaRule("user@Example.Com"))
}

func TestMatchEmailQuotaRule_DisabledRule(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{
EmailSuffix: "@disabled.com",
Quota: 100000,
Enabled: false,
}
require.NoError(t, rule.Insert())

result := MatchEmailQuotaRule("user@disabled.com")
assert.Equal(t, int64(-1), result)
}

func TestMatchEmailQuotaRule_NoMatch(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{
EmailSuffix: "@company.com",
Quota: 500000,
Enabled: true,
}
require.NoError(t, rule.Insert())

result := MatchEmailQuotaRule("user@other.com")
assert.Equal(t, int64(-1), result)
}

func TestEmailQuotaRule_Insert(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{
EmailSuffix: "@test.com",
Quota: 100000,
Enabled: true,
Description: "Test rule",
}
require.NoError(t, rule.Insert())
assert.Greater(t, rule.Id, 0)
assert.Greater(t, rule.CreatedTime, int64(0))
assert.Equal(t, rule.CreatedTime, rule.UpdatedTime)

// Verify cache is populated
assert.Equal(t, int64(100000), MatchEmailQuotaRule("user@test.com"))
}

func TestEmailQuotaRule_Insert_DuplicateSuffix(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule1 := &EmailQuotaRule{EmailSuffix: "@dup.com", Quota: 100, Enabled: true}
require.NoError(t, rule1.Insert())

rule2 := &EmailQuotaRule{EmailSuffix: "@dup.com", Quota: 200, Enabled: true}
assert.Error(t, rule2.Insert())
}

func TestEmailQuotaRule_Update(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{EmailSuffix: "@update.com", Quota: 100, Enabled: true}
require.NoError(t, rule.Insert())

rule.Quota = 999
rule.Description = "updated"
require.NoError(t, rule.Update())

// Verify DB
found, err := GetEmailQuotaRuleById(rule.Id)
require.NoError(t, err)
assert.Equal(t, int64(999), found.Quota)
assert.Equal(t, "updated", found.Description)

// Verify cache refreshed
assert.Equal(t, int64(999), MatchEmailQuotaRule("user@update.com"))
}

func TestEmailQuotaRule_Update_Disable(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{EmailSuffix: "@toggled.com", Quota: 500, Enabled: true}
require.NoError(t, rule.Insert())
assert.Equal(t, int64(500), MatchEmailQuotaRule("user@toggled.com"))

rule.Enabled = false
require.NoError(t, rule.Update())

assert.Equal(t, int64(-1), MatchEmailQuotaRule("user@toggled.com"))
}

func TestEmailQuotaRule_Delete(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{EmailSuffix: "@delete.com", Quota: 100, Enabled: true}
require.NoError(t, rule.Insert())
assert.Equal(t, int64(100), MatchEmailQuotaRule("user@delete.com"))

require.NoError(t, rule.Delete())
assert.Equal(t, int64(-1), MatchEmailQuotaRule("user@delete.com"))
}

func TestGetAllEmailQuotaRules(t *testing.T) {
setupEmailQuotaRuleDB(t)

r1 := &EmailQuotaRule{EmailSuffix: "@a.com", Quota: 100, Enabled: true}
r2 := &EmailQuotaRule{EmailSuffix: "@b.com", Quota: 200, Enabled: false}
require.NoError(t, r1.Insert())
require.NoError(t, r2.Insert())

list, err := GetAllEmailQuotaRules()
require.NoError(t, err)
assert.Len(t, list, 2)
// Ordered by id ASC
assert.Equal(t, "@a.com", list[0].EmailSuffix)
assert.Equal(t, "@b.com", list[1].EmailSuffix)
}

func TestGetEmailQuotaRuleBySuffix(t *testing.T) {
setupEmailQuotaRuleDB(t)

rule := &EmailQuotaRule{EmailSuffix: "@find.com", Quota: 300, Enabled: true}
require.NoError(t, rule.Insert())

found, err := GetEmailQuotaRuleBySuffix("@find.com")
require.NoError(t, err)
assert.Equal(t, int64(300), found.Quota)

_, err = GetEmailQuotaRuleBySuffix("@notexist.com")
assert.Error(t, err)
}

func TestLoadEmailQuotaCache_OnlyEnabled(t *testing.T) {
setupEmailQuotaRuleDB(t)

r1 := &EmailQuotaRule{EmailSuffix: "@enabled.com", Quota: 100, Enabled: true}
r2 := &EmailQuotaRule{EmailSuffix: "@disabled.com", Quota: 200, Enabled: false}
require.NoError(t, r1.Insert())
require.NoError(t, r2.Insert())
LoadEmailQuotaCache()

assert.Equal(t, int64(100), MatchEmailQuotaRule("user@enabled.com"))
assert.Equal(t, int64(-1), MatchEmailQuotaRule("user@disabled.com"))
}

func TestMatchEmailQuotaRule_MultipleRules(t *testing.T) {
setupEmailQuotaRuleDB(t)

rules := []*EmailQuotaRule{
{EmailSuffix: "@company.com", Quota: 500000, Enabled: true},
{EmailSuffix: "@tsinghua.edu.cn", Quota: 1000000, Enabled: true},
{EmailSuffix: "@vip.org", Quota: 2000000, Enabled: true},
}
for _, r := range rules {
require.NoError(t, r.Insert())
}

assert.Equal(t, int64(500000), MatchEmailQuotaRule("user@company.com"))
assert.Equal(t, int64(1000000), MatchEmailQuotaRule("student@tsinghua.edu.cn"))
assert.Equal(t, int64(2000000), MatchEmailQuotaRule("admin@vip.org"))
assert.Equal(t, int64(-1), MatchEmailQuotaRule("random@unknown.net"))
}

+ 26
- 3
model/main.go Переглянути файл

@@ -203,7 +203,12 @@ func InitDB() (err error) {
} }
common.SysLog("database migration started") common.SysLog("database migration started")
err = migrateDB() err = migrateDB()
return err
if err != nil {
return err
}
LoadEmailQuotaCache()
LoadChannelPricingCache()
return nil
} else { } else {
common.FatalLog(err) common.FatalLog(err)
} }
@@ -282,6 +287,7 @@ func migrateDB() error {
&PricingTag{}, &PricingTag{},
&PendingSyncRecord{}, &PendingSyncRecord{},
&QuotaSyncLog{}, &QuotaSyncLog{},
&EmailQuotaRule{},
) )
if err != nil { if err != nil {
return err return err
@@ -294,9 +300,14 @@ func migrateDB() error {
if err := DB.AutoMigrate(&SubscriptionPlan{}); err != nil { if err := DB.AutoMigrate(&SubscriptionPlan{}); err != nil {
return err return err
} }
}
// 将现有 sort_order=0 的模型和供应商更新为默认大数
DB.Model(&Model{}).Where("sort_order = 0").Update("sort_order", 999999)
DB.Model(&Vendor{}).Where("sort_order = 0").Update("sort_order", 999999)
migrateChannelPublicName()

return nil
} }
return nil
}


func migrateDBFast() error { func migrateDBFast() error {
// Drop bound_channel_id column from tokens table (deprecated field) // Drop bound_channel_id column from tokens table (deprecated field)
@@ -336,6 +347,7 @@ func migrateDBFast() error {
{&PricingTag{}, "PricingTag"}, {&PricingTag{}, "PricingTag"},
{&PendingSyncRecord{}, "PendingSyncRecord"}, {&PendingSyncRecord{}, "PendingSyncRecord"},
{&QuotaSyncLog{}, "QuotaSyncLog"}, {&QuotaSyncLog{}, "QuotaSyncLog"},
{&EmailQuotaRule{}, "EmailQuotaRule"},
} }
// 动态计算migration数量,确保errChan缓冲区足够大 // 动态计算migration数量,确保errChan缓冲区足够大
errChan := make(chan error, len(migrations)) errChan := make(chan error, len(migrations))
@@ -683,3 +695,14 @@ func PingDB() error {
common.SysLog("Database pinged successfully") common.SysLog("Database pinged successfully")
return nil return nil
} }

func migrateChannelPublicName() {
result := DB.Model(&Channel{}).
Where("public_name = '' OR public_name IS NULL").
Update("public_name", gorm.Expr("name"))
if result.Error != nil {
common.SysError("[Migration] migrateChannelPublicName failed: " + result.Error.Error())
} else if result.RowsAffected > 0 {
common.SysLog(fmt.Sprintf("[Migration] migrateChannelPublicName: backfilled %d channels", result.RowsAffected))
}
}

+ 29
- 3
model/model_meta.go Переглянути файл

@@ -42,6 +42,7 @@ type Model struct {
Endpoints string `json:"endpoints,omitempty" gorm:"type:text"` Endpoints string `json:"endpoints,omitempty" gorm:"type:text"`
Status int `json:"status" gorm:"default:1"` Status int `json:"status" gorm:"default:1"`
SyncOfficial int `json:"sync_official" gorm:"default:1"` SyncOfficial int `json:"sync_official" gorm:"default:1"`
SortOrder int `json:"sort_order" gorm:"default:999999"`
CreatedTime int64 `json:"created_time" gorm:"bigint"` CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"` UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
DeletedAt gorm.DeletedAt `json:"-" gorm:"index;uniqueIndex:uk_model_name_delete_at,priority:2"` DeletedAt gorm.DeletedAt `json:"-" gorm:"index;uniqueIndex:uk_model_name_delete_at,priority:2"`
@@ -89,7 +90,7 @@ func (mi *Model) Update() error {
mi.UpdatedTime = common.GetTimestamp() mi.UpdatedTime = common.GetTimestamp()
// 使用 Select 强制更新所有字段,包括零值 // 使用 Select 强制更新所有字段,包括零值
return DB.Model(&Model{}).Where("id = ?", mi.Id). return DB.Model(&Model{}).Where("id = ?", mi.Id).
Select("model_name", "description", "icon", "tags", "type", "vendor_id", "endpoints", "status", "sync_official", "name_rule", "updated_time").
Select("model_name", "description", "icon", "tags", "type", "vendor_id", "endpoints", "status", "sync_official", "name_rule", "sort_order", "updated_time").
Updates(mi).Error Updates(mi).Error
} }


@@ -97,6 +98,16 @@ func (mi *Model) Delete() error {
return DB.Delete(mi).Error return DB.Delete(mi).Error
} }


// GetDisabledModelNames returns model names that are disabled (status = 0) from the given list.
func GetDisabledModelNames(names []string) []string {
if len(names) == 0 {
return nil
}
var disabled []string
DB.Table("models").Where("model_name IN ? AND status = 0", names).Pluck("model_name", &disabled)
return disabled
}

func GetVendorModelCounts() (map[int64]int64, error) { func GetVendorModelCounts() (map[int64]int64, error) {
var stats []struct { var stats []struct {
VendorID int64 VendorID int64
@@ -117,7 +128,7 @@ func GetVendorModelCounts() (map[int64]int64, error) {


func GetAllModels(offset int, limit int) ([]*Model, error) { func GetAllModels(offset int, limit int) ([]*Model, error) {
var models []*Model var models []*Model
err := DB.Order("id DESC").Offset(offset).Limit(limit).Find(&models).Error
err := DB.Order("sort_order ASC, id ASC").Offset(offset).Limit(limit).Find(&models).Error
return models, err return models, err
} }


@@ -165,8 +176,23 @@ func SearchModels(keyword string, vendor string, offset int, limit int) ([]*Mode
if err := db.Count(&total).Error; err != nil { if err := db.Count(&total).Error; err != nil {
return nil, 0, err return nil, 0, err
} }
if err := db.Order("models.id DESC").Offset(offset).Limit(limit).Find(&models).Error; err != nil {
if err := db.Order("sort_order ASC, models.id ASC").Offset(offset).Limit(limit).Find(&models).Error; err != nil {
return nil, 0, err return nil, 0, err
} }
return models, total, nil return models, total, nil
} }

// ReorderModels 批量更新模型排序值
func ReorderModels(items []struct {
Id int `json:"id"`
SortOrder int `json:"sort_order"`
}) error {
return DB.Transaction(func(tx *gorm.DB) error {
for _, item := range items {
if err := tx.Model(&Model{}).Where("id = ?", item.Id).Update("sort_order", item.SortOrder).Error; err != nil {
return err
}
}
return nil
})
}

+ 33
- 0
model/option.go Переглянути файл

@@ -43,6 +43,7 @@ func InitOptionMap() {
common.OptionMap["TelegramOAuthEnabled"] = strconv.FormatBool(common.TelegramOAuthEnabled) common.OptionMap["TelegramOAuthEnabled"] = strconv.FormatBool(common.TelegramOAuthEnabled)
common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled) common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled)
common.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(common.TurnstileCheckEnabled) common.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(common.TurnstileCheckEnabled)
common.OptionMap["CaptchaEnabled"] = strconv.FormatBool(common.CaptchaEnabled)
common.OptionMap["RegisterEnabled"] = strconv.FormatBool(common.RegisterEnabled) common.OptionMap["RegisterEnabled"] = strconv.FormatBool(common.RegisterEnabled)
common.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(common.AutomaticDisableChannelEnabled) common.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(common.AutomaticDisableChannelEnabled)
common.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(common.AutomaticEnableChannelEnabled) common.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(common.AutomaticEnableChannelEnabled)
@@ -101,6 +102,12 @@ func InitOptionMap() {
common.OptionMap["WechatPayPubKeyB64"] = setting.WechatPayPubKeyB64 common.OptionMap["WechatPayPubKeyB64"] = setting.WechatPayPubKeyB64
common.OptionMap["WechatPayMinTopUp"] = strconv.Itoa(setting.WechatPayMinTopUp) common.OptionMap["WechatPayMinTopUp"] = strconv.Itoa(setting.WechatPayMinTopUp)
common.OptionMap["WechatPayUnitPrice"] = strconv.FormatFloat(setting.WechatPayUnitPrice, 'f', -1, 64) common.OptionMap["WechatPayUnitPrice"] = strconv.FormatFloat(setting.WechatPayUnitPrice, 'f', -1, 64)
common.OptionMap["AlipayAppID"] = setting.AlipayAppID
common.OptionMap["AlipayPrivateKey"] = setting.AlipayPrivateKey
common.OptionMap["AlipayPublicKey"] = setting.AlipayPublicKey
common.OptionMap["AlipayNotifyURL"] = setting.AlipayNotifyURL
common.OptionMap["AlipayMinTopUp"] = strconv.Itoa(setting.AlipayMinTopUp)
common.OptionMap["AlipayUnitPrice"] = strconv.FormatFloat(setting.AlipayUnitPrice, 'f', -1, 64)
common.OptionMap["TopupGroupRatio"] = common.TopupGroupRatio2JSONString() common.OptionMap["TopupGroupRatio"] = common.TopupGroupRatio2JSONString()
common.OptionMap["Chats"] = setting.Chats2JsonString() common.OptionMap["Chats"] = setting.Chats2JsonString()
common.OptionMap["AutoGroups"] = setting.AutoGroups2JsonString() common.OptionMap["AutoGroups"] = setting.AutoGroups2JsonString()
@@ -178,6 +185,13 @@ func triggerWechatPayReset() {
} }
} }


// triggerAlipayReset 安全触发支付宝客户端重置
func triggerAlipayReset() {
if setting.OnAlipayConfigChanged != nil {
setting.OnAlipayConfigChanged()
}
}

func loadOptionsFromDatabase() { func loadOptionsFromDatabase() {
options, _ := AllOption() options, _ := AllOption()
for _, option := range options { for _, option := range options {
@@ -255,6 +269,8 @@ func updateOptionMap(key string, value string) (err error) {
common.TelegramOAuthEnabled = boolValue common.TelegramOAuthEnabled = boolValue
case "TurnstileCheckEnabled": case "TurnstileCheckEnabled":
common.TurnstileCheckEnabled = boolValue common.TurnstileCheckEnabled = boolValue
case "CaptchaEnabled":
common.CaptchaEnabled = boolValue
case "RegisterEnabled": case "RegisterEnabled":
common.RegisterEnabled = boolValue common.RegisterEnabled = boolValue
case "EmailDomainRestrictionEnabled": case "EmailDomainRestrictionEnabled":
@@ -410,6 +426,21 @@ func updateOptionMap(key string, value string) (err error) {
setting.WechatPayMinTopUp, _ = strconv.Atoi(value) setting.WechatPayMinTopUp, _ = strconv.Atoi(value)
case "WechatPayUnitPrice": case "WechatPayUnitPrice":
setting.WechatPayUnitPrice, _ = strconv.ParseFloat(value, 64) setting.WechatPayUnitPrice, _ = strconv.ParseFloat(value, 64)
case "AlipayAppID":
setting.AlipayAppID = value
triggerAlipayReset()
case "AlipayPrivateKey":
setting.AlipayPrivateKey = value
triggerAlipayReset()
case "AlipayPublicKey":
setting.AlipayPublicKey = value
triggerAlipayReset()
case "AlipayNotifyURL":
setting.AlipayNotifyURL = value
case "AlipayMinTopUp":
setting.AlipayMinTopUp, _ = strconv.Atoi(value)
case "AlipayUnitPrice":
setting.AlipayUnitPrice, _ = strconv.ParseFloat(value, 64)
case "TopupGroupRatio": case "TopupGroupRatio":
err = common.UpdateTopupGroupRatioByJSONString(value) err = common.UpdateTopupGroupRatioByJSONString(value)
case "GitHubClientId": case "GitHubClientId":
@@ -428,6 +459,8 @@ func updateOptionMap(key string, value string) (err error) {
common.SystemName = value common.SystemName = value
case "Logo": case "Logo":
common.Logo = value common.Logo = value
case "DefaultLanguage":
common.DefaultLanguage = value
case "WeChatServerAddress": case "WeChatServerAddress":
common.WeChatServerAddress = value common.WeChatServerAddress = value
case "WeChatServerToken": case "WeChatServerToken":


+ 134
- 9
model/pricing.go Переглянути файл

@@ -3,6 +3,7 @@ package model
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"sort"
"strings" "strings"


"sync" "sync"
@@ -25,10 +26,13 @@ type Pricing struct {
ModelPrice float64 `json:"model_price"` ModelPrice float64 `json:"model_price"`
OwnerBy string `json:"owner_by"` OwnerBy string `json:"owner_by"`
CompletionRatio float64 `json:"completion_ratio"` CompletionRatio float64 `json:"completion_ratio"`
CacheRatio float64 `json:"cache_ratio"`
CacheCreationRatio float64 `json:"cache_creation_ratio"`
EnableGroup []string `json:"enable_groups"` EnableGroup []string `json:"enable_groups"`
SupportedEndpointTypes []constant.EndpointType `json:"supported_endpoint_types"` SupportedEndpointTypes []constant.EndpointType `json:"supported_endpoint_types"`
PricingVersion string `json:"pricing_version,omitempty"` PricingVersion string `json:"pricing_version,omitempty"`
Type int `json:"type"` Type int `json:"type"`
DefaultChannelName string `json:"default_channel_name,omitempty"`
} }


type PricingVendor struct { type PricingVendor struct {
@@ -162,6 +166,11 @@ func updatePricing() {
initDefaultVendorMapping(metaMap, vendorMap, enableAbilities) initDefaultVendorMapping(metaMap, vendorMap, enableAbilities)


// 构建对前端友好的供应商列表 // 构建对前端友好的供应商列表
vendorOrderMap := make(map[int]int)
for _, v := range vendorMap {
vendorOrderMap[v.Id] = v.SortOrder
}

vendorsList = make([]PricingVendor, 0, len(vendorMap)) vendorsList = make([]PricingVendor, 0, len(vendorMap))
for _, v := range vendorMap { for _, v := range vendorMap {
vendorsList = append(vendorsList, PricingVendor{ vendorsList = append(vendorsList, PricingVendor{
@@ -171,6 +180,14 @@ func updatePricing() {
Icon: v.Icon, Icon: v.Icon,
}) })
} }
sort.Slice(vendorsList, func(i, j int) bool {
oi := vendorOrderMap[vendorsList[i].ID]
oj := vendorOrderMap[vendorsList[j].ID]
if oi != oj {
return oi < oj
}
return vendorsList[i].ID < vendorsList[j].ID
})


modelGroupsMap := make(map[string]*types.Set[string]) modelGroupsMap := make(map[string]*types.Set[string])


@@ -269,6 +286,24 @@ func updatePricing() {
} }
} }


// 从渠道定价表加载实际定价数据(仅启用渠道),同时获取渠道名称
var allCPs []struct {
ChannelPricing
ChannelName string
ChannelPublicName string
}
DB.Table("channel_pricings").
Select("channel_pricings.*, channels.name as channel_name, channels.public_name as channel_public_name").
Joins("JOIN channels ON channel_pricings.channel_id = channels.id").
Where("channels.status = 1 AND channel_pricings.deleted_at IS NULL").
Find(&allCPs)
cpMap := make(map[string][]ChannelPricing)
channelNameMap := make(map[int]string)
for i := range allCPs {
cpMap[allCPs[i].ModelName] = append(cpMap[allCPs[i].ModelName], allCPs[i].ChannelPricing)
channelNameMap[allCPs[i].ChannelId] = ChannelDisplayName(allCPs[i].ChannelPublicName, allCPs[i].ChannelName)
}

pricingMap = make([]Pricing, 0) pricingMap = make([]Pricing, 0)
for model, groups := range modelGroupsMap { for model, groups := range modelGroupsMap {
pricing := Pricing{ pricing := Pricing{
@@ -289,19 +324,37 @@ func updatePricing() {
pricing.VendorID = meta.VendorID pricing.VendorID = meta.VendorID
pricing.Type = meta.Type pricing.Type = meta.Type
} }
modelPrice, findPrice := ratio_setting.GetModelPrice(model, false)
if findPrice {
pricing.ModelPrice = modelPrice
pricing.QuotaType = 1
} else {
modelRatio, _, _ := ratio_setting.GetModelRatio(model)
pricing.ModelRatio = modelRatio
pricing.CompletionRatio = ratio_setting.GetCompletionRatio(model)
pricing.QuotaType = 0

// 使用渠道定价表中的实际数据,选取最便宜的渠道
applyBestChannelPricing(&pricing, cpMap[model], model)
// 填充默认通道名称
if chId, ok := GetDefaultChannelId(model); ok {
if name, found := channelNameMap[chId]; found {
pricing.DefaultChannelName = name
}
} }

pricingMap = append(pricingMap, pricing) pricingMap = append(pricingMap, pricing)
} }


// 按 sort_order 排序 pricingMap,999999 视为未设置
sort.Slice(pricingMap, func(i, j int) bool {
mi, okI := metaMap[pricingMap[i].ModelName]
mj, okJ := metaMap[pricingMap[j].ModelName]
si := 999999
sj := 999999
if okI {
si = mi.SortOrder
}
if okJ {
sj = mj.SortOrder
}
if si != sj {
return si < sj
}
return pricingMap[i].ModelName < pricingMap[j].ModelName
})

// 防止大更新后数据不通用 // 防止大更新后数据不通用
if len(pricingMap) > 0 { if len(pricingMap) > 0 {
pricingMap[0].PricingVersion = "82c4a357505fff6fee8462c3f7ec8a645bb95532669cb73b2cabee6a416ec24f" pricingMap[0].PricingVersion = "82c4a357505fff6fee8462c3f7ec8a645bb95532669cb73b2cabee6a416ec24f"
@@ -324,3 +377,75 @@ func updatePricing() {
func GetSupportedEndpointMap() map[string]common.EndpointInfo { func GetSupportedEndpointMap() map[string]common.EndpointInfo {
return supportedEndpointMap return supportedEndpointMap
} }

// applyGlobalDefault 用全局默认值填充 Pricing(无渠道定价时的回退)
func applyGlobalDefault(pricing *Pricing, model string) {
modelPrice, findPrice := ratio_setting.GetModelPrice(model, false)
if findPrice {
pricing.ModelPrice = modelPrice
pricing.QuotaType = 1
} else {
modelRatio, _, _ := ratio_setting.GetModelRatio(model)
pricing.ModelRatio = modelRatio
pricing.CompletionRatio = ratio_setting.GetCompletionRatio(model)
pricing.QuotaType = 0
}
pricing.CacheRatio, _ = ratio_setting.GetCacheRatio(model)
pricing.CacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(model)
}

// applyBestChannelPricing 从渠道定价中选取最优(最便宜)的价格填充 Pricing
// 优先选按量计费(quota_type=0)的渠道,因为缓存价格仅对按量计费有意义
// 扩展比率字段(cache_ratio 等)为 0 表示未设置,需回退到全局默认值
func applyBestChannelPricing(pricing *Pricing, cps []ChannelPricing, model string) {
if len(cps) == 0 {
applyGlobalDefault(pricing, model)
return
}

// 优先选按量计费 (quota_type=0) 中 model_ratio 最低的渠道
var bestPerToken *ChannelPricing
for i := range cps {
cp := &cps[i]
if cp.QuotaType == 0 {
if bestPerToken == nil || cp.ModelRatio < bestPerToken.ModelRatio {
bestPerToken = cp
}
}
}

if bestPerToken != nil {
pricing.QuotaType = 0
pricing.ModelRatio = bestPerToken.ModelRatio
pricing.CompletionRatio = bestPerToken.CompletionRatio
if bestPerToken.CacheRatio > 0 {
pricing.CacheRatio = bestPerToken.CacheRatio
} else {
pricing.CacheRatio, _ = ratio_setting.GetCacheRatio(model)
}
if bestPerToken.CacheCreationRatio > 0 {
pricing.CacheCreationRatio = bestPerToken.CacheCreationRatio
} else {
pricing.CacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(model)
}
return
}

// 没有按量渠道,选按次计费 (quota_type=1) 中 model_price 最低的渠道
var bestPerCall *ChannelPricing
for i := range cps {
cp := &cps[i]
if cp.QuotaType == 1 {
if bestPerCall == nil || cp.ModelPrice < bestPerCall.ModelPrice {
bestPerCall = cp
}
}
}
if bestPerCall != nil {
pricing.QuotaType = 1
pricing.ModelPrice = bestPerCall.ModelPrice
return
}

applyGlobalDefault(pricing, model)
}

+ 122
- 0
model/pricing_test.go Переглянути файл

@@ -0,0 +1,122 @@
package model

import (
"testing"

"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

const testPricingModel = "test-pricing-model-apply"

func setupPricingTest(t *testing.T) {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, _ := db.DB()
sqlDB.SetMaxOpenConns(1)

origDB := DB
DB = db
require.NoError(t, db.AutoMigrate(&ChannelPricing{}))

// 全局默认定价
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(`{"`+testPricingModel+`":10}`))
require.NoError(t, ratio_setting.UpdateCompletionRatioByJSONString(`{"`+testPricingModel+`":3}`))
require.NoError(t, ratio_setting.UpdateCacheRatioByJSONString(`{"`+testPricingModel+`":0.5}`))
require.NoError(t, ratio_setting.UpdateCreateCacheRatioByJSONString(`{"`+testPricingModel+`":0.75}`))

t.Cleanup(func() {
DB = origDB
sqlDB.Close()
ratio_setting.UpdateModelRatioByJSONString(`{}`)
ratio_setting.UpdateCompletionRatioByJSONString(`{}`)
ratio_setting.UpdateCacheRatioByJSONString(`{}`)
ratio_setting.UpdateCreateCacheRatioByJSONString(`{}`)
ratio_setting.UpdateModelPriceByJSONString(`{}`)
})
}

func TestApplyBestChannelPricing_NoChannelPricing(t *testing.T) {
setupPricingTest(t)

p := &Pricing{}
applyBestChannelPricing(p, nil, testPricingModel)

require.Equal(t, 0, p.QuotaType, "应使用全局按量计费")
require.Equal(t, 10.0, p.ModelRatio, "应使用全局 model_ratio")
require.Equal(t, 3.0, p.CompletionRatio, "应使用全局 completion_ratio")
require.Equal(t, 0.5, p.CacheRatio, "应使用全局 cache_ratio")
require.Equal(t, 0.75, p.CacheCreationRatio, "应使用全局 cache_creation_ratio")
}

func TestApplyBestChannelPricing_PerTokenCheapest(t *testing.T) {
setupPricingTest(t)

// 渠道A: model_ratio=8(更便宜)
require.NoError(t, (&ChannelPricing{
ModelName: testPricingModel, ChannelId: 1,
QuotaType: QuotaTypeByTokens, ModelRatio: 8, CompletionRatio: 2,
CacheRatio: 0.3, CacheCreationRatio: 0.6,
}).Insert())
// 渠道B: model_ratio=12(更贵)
require.NoError(t, (&ChannelPricing{
ModelName: testPricingModel, ChannelId: 2,
QuotaType: QuotaTypeByTokens, ModelRatio: 12, CompletionRatio: 4,
CacheRatio: 0.8, CacheCreationRatio: 1.0,
}).Insert())

cps := []ChannelPricing{
{ModelRatio: 8, CompletionRatio: 2, CacheRatio: 0.3, CacheCreationRatio: 0.6},
{ModelRatio: 12, CompletionRatio: 4, CacheRatio: 0.8, CacheCreationRatio: 1.0},
}

p := &Pricing{}
applyBestChannelPricing(p, cps, testPricingModel)

require.Equal(t, 0, p.QuotaType)
require.Equal(t, 8.0, p.ModelRatio, "应选最便宜的渠道A")
require.Equal(t, 2.0, p.CompletionRatio)
require.Equal(t, 0.3, p.CacheRatio)
require.Equal(t, 0.6, p.CacheCreationRatio)
}

func TestApplyBestChannelPricing_ExtendedRatioZeroFallback(t *testing.T) {
setupPricingTest(t)

// 渠道定价中扩展比率为 0,应回退到全局值
require.NoError(t, (&ChannelPricing{
ModelName: testPricingModel, ChannelId: 1,
QuotaType: QuotaTypeByTokens, ModelRatio: 5, CompletionRatio: 1,
CacheRatio: 0, CacheCreationRatio: 0, // 0 = 未设置
}).Insert())

cps := []ChannelPricing{
{ModelRatio: 5, CompletionRatio: 1, CacheRatio: 0, CacheCreationRatio: 0},
}

p := &Pricing{}
applyBestChannelPricing(p, cps, testPricingModel)

require.Equal(t, 0.5, p.CacheRatio, "cache_ratio=0 应回退全局 0.5")
require.Equal(t, 0.75, p.CacheCreationRatio, "cache_creation_ratio=0 应回退全局 0.75")
}

func TestApplyBestChannelPricing_PerCallCheapest(t *testing.T) {
setupPricingTest(t)
require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{"`+testPricingModel+`":0.8}`))

// 只有按次计费的渠道
cps := []ChannelPricing{
{QuotaType: QuotaTypeByCall, ModelPrice: 0.3},
{QuotaType: QuotaTypeByCall, ModelPrice: 0.5},
}

p := &Pricing{}
applyBestChannelPricing(p, cps, testPricingModel)

require.Equal(t, 1, p.QuotaType)
require.Equal(t, 0.3, p.ModelPrice, "应选最便宜的按次渠道")
}

+ 4
- 3
model/redemption.go Переглянути файл

@@ -20,6 +20,7 @@ type Redemption struct {
Key string `json:"key" gorm:"type:char(32);uniqueIndex"` Key string `json:"key" gorm:"type:char(32);uniqueIndex"`
Status int `json:"status" gorm:"default:1"` Status int `json:"status" gorm:"default:1"`
Name string `json:"name" gorm:"index"` Name string `json:"name" gorm:"index"`
Remark string `json:"remark" gorm:"index"`
Quota int `json:"quota" gorm:"default:100"` Quota int `json:"quota" gorm:"default:100"`
CreatedTime int64 `json:"created_time" gorm:"bigint"` CreatedTime int64 `json:"created_time" gorm:"bigint"`
RedeemedTime int64 `json:"redeemed_time" gorm:"bigint"` RedeemedTime int64 `json:"redeemed_time" gorm:"bigint"`
@@ -79,9 +80,9 @@ func SearchRedemptions(keyword string, startIdx int, num int) (redemptions []*Re


// Only try to convert to ID if the string represents a valid integer // Only try to convert to ID if the string represents a valid integer
if id, err := strconv.Atoi(keyword); err == nil { if id, err := strconv.Atoi(keyword); err == nil {
query = query.Where("id = ? OR name LIKE ?", id, keyword+"%")
query = query.Where("id = ? OR name LIKE ? OR remark LIKE ?", id, keyword+"%", keyword+"%")
} else { } else {
query = query.Where("name LIKE ?", keyword+"%")
query = query.Where("name LIKE ? OR remark LIKE ?", keyword+"%", keyword+"%")
} }


// Get total count // Get total count
@@ -187,7 +188,7 @@ func (redemption *Redemption) SelectUpdate() error {
// Update Make sure your token's fields is completed, because this will update non-zero values // Update Make sure your token's fields is completed, because this will update non-zero values
func (redemption *Redemption) Update() error { func (redemption *Redemption) Update() error {
var err error var err error
err = DB.Model(redemption).Select("name", "status", "quota", "redeemed_time", "expired_time").Updates(redemption).Error
err = DB.Model(redemption).Select("name", "remark", "status", "quota", "redeemed_time", "expired_time").Updates(redemption).Error
return err return err
} }




+ 101
- 14
model/redemption_test.go Переглянути файл

@@ -48,13 +48,13 @@ func TestRedeem_Success(t *testing.T) {


// 创建兑换码 // 创建兑换码
redemption := Redemption{ redemption := Redemption{
Id: 1,
UserId: 1,
Key: "test-key-123",
Status: common.RedemptionCodeStatusEnabled,
Quota: 50000,
CreatedTime: common.GetTimestamp(),
ExpiredTime: 0, // 永不过期
Id: 1,
UserId: 1,
Key: "test-key-123",
Status: common.RedemptionCodeStatusEnabled,
Quota: 50000,
CreatedTime: common.GetTimestamp(),
ExpiredTime: 0, // 永不过期
} }
require.NoError(t, db.Create(&redemption).Error) require.NoError(t, db.Create(&redemption).Error)


@@ -148,13 +148,13 @@ func TestRedeem_Expired(t *testing.T) {


// 创建已过期的兑换码 // 创建已过期的兑换码
redemption := Redemption{ redemption := Redemption{
Id: 1,
UserId: 1,
Key: "expired-key",
Status: common.RedemptionCodeStatusEnabled,
Quota: 50000,
CreatedTime: common.GetTimestamp() - 86400,
ExpiredTime: common.GetTimestamp() - 3600, // 已过期
Id: 1,
UserId: 1,
Key: "expired-key",
Status: common.RedemptionCodeStatusEnabled,
Quota: 50000,
CreatedTime: common.GetTimestamp() - 86400,
ExpiredTime: common.GetTimestamp() - 3600, // 已过期
} }
require.NoError(t, db.Create(&redemption).Error) require.NoError(t, db.Create(&redemption).Error)


@@ -252,3 +252,90 @@ func TestRedeem_SyncedUser(t *testing.T) {
require.NoError(t, db.First(&updatedUser, 100).Error) require.NoError(t, db.First(&updatedUser, 100).Error)
assert.Equal(t, 150000, updatedUser.Quota) assert.Equal(t, 150000, updatedUser.Quota)
} }

func TestRedemptionUpdateRemarkPersistsRemark(t *testing.T) {
db := setupRedemptionDB(t)

redemption := Redemption{
Id: 1,
UserId: 1,
Key: "remark-update-key",
Status: common.RedemptionCodeStatusEnabled,
Name: "starter",
Remark: "initial-note",
Quota: 100,
CreatedTime: common.GetTimestamp(),
}
require.NoError(t, db.Create(&redemption).Error)

redemption.Remark = "vip-updated"
require.NoError(t, redemption.Update())

var updated Redemption
require.NoError(t, db.First(&updated, redemption.Id).Error)
assert.Equal(t, "vip-updated", updated.Remark)
}

func TestSearchRedemptionsByRemark(t *testing.T) {
db := setupRedemptionDB(t)

require.NoError(t, db.Create(&Redemption{
Id: 1,
UserId: 1,
Key: "remark-search-key",
Status: common.RedemptionCodeStatusEnabled,
Name: "starter-pack",
Remark: "vip benefit",
Quota: 100,
CreatedTime: common.GetTimestamp(),
}).Error)
require.NoError(t, db.Create(&Redemption{
Id: 2,
UserId: 1,
Key: "remark-search-other",
Status: common.RedemptionCodeStatusEnabled,
Name: "basic-pack",
Remark: "standard benefit",
Quota: 100,
CreatedTime: common.GetTimestamp(),
}).Error)

results, total, err := SearchRedemptions("vip", 0, 10)
require.NoError(t, err)
require.Equal(t, int64(1), total)
require.Len(t, results, 1)
assert.Equal(t, 1, results[0].Id)
assert.Equal(t, "vip benefit", results[0].Remark)
}

func TestSearchRedemptionsByRemarkWithNumericKeyword(t *testing.T) {
db := setupRedemptionDB(t)

require.NoError(t, db.Create(&Redemption{
Id: 1,
UserId: 1,
Key: "remark-search-numeric",
Status: common.RedemptionCodeStatusEnabled,
Name: "starter-pack",
Remark: "123-vip benefit",
Quota: 100,
CreatedTime: common.GetTimestamp(),
}).Error)
require.NoError(t, db.Create(&Redemption{
Id: 2,
UserId: 1,
Key: "remark-search-other-numeric",
Status: common.RedemptionCodeStatusEnabled,
Name: "basic-pack",
Remark: "standard benefit",
Quota: 100,
CreatedTime: common.GetTimestamp(),
}).Error)

results, total, err := SearchRedemptions("123", 0, 10)
require.NoError(t, err)
require.Equal(t, int64(1), total)
require.Len(t, results, 1)
assert.Equal(t, 1, results[0].Id)
assert.Equal(t, "123-vip benefit", results[0].Remark)
}

+ 32
- 1
model/topup.go Переглянути файл

@@ -5,6 +5,7 @@ import (
"fmt" "fmt"


"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/types"
"github.com/QuantumNous/new-api/logger" "github.com/QuantumNous/new-api/logger"


"github.com/shopspring/decimal" "github.com/shopspring/decimal"
@@ -21,6 +22,32 @@ type TopUp struct {
CreateTime int64 `json:"create_time"` CreateTime int64 `json:"create_time"`
CompleteTime int64 `json:"complete_time"` CompleteTime int64 `json:"complete_time"`
Status string `json:"status"` Status string `json:"status"`
UserEmail string `json:"user_email" gorm:"-"` // Join 查询时填充,非数据库字段
}

// fillTopUpEmails 批量填充 topup 记录的用户邮箱
func fillTopUpEmails(topups []*TopUp) {
if len(topups) == 0 {
return
}
userIds := types.NewSet[int]()
for _, t := range topups {
userIds.Add(t.UserId)
}

var users []User
if err := DB.Select("id, email").Where("id IN ?", userIds.Items()).Find(&users).Error; err != nil {
common.SysError("fillTopUpEmails: " + err.Error())
return
}

emailMap := make(map[int]string, len(users))
for _, u := range users {
emailMap[u.Id] = u.Email
}
for _, t := range topups {
t.UserEmail = emailMap[t.UserId]
}
} }


func (topUp *TopUp) Insert() error { func (topUp *TopUp) Insert() error {
@@ -135,6 +162,7 @@ func GetUserTopUps(userId int, pageInfo *common.PageInfo) (topups []*TopUp, tota
return nil, 0, err return nil, 0, err
} }


fillTopUpEmails(topups)
return topups, total, nil return topups, total, nil
} }


@@ -164,6 +192,7 @@ func GetAllTopUps(pageInfo *common.PageInfo) (topups []*TopUp, total int64, err
return nil, 0, err return nil, 0, err
} }


fillTopUpEmails(topups)
return topups, total, nil return topups, total, nil
} }


@@ -198,6 +227,7 @@ func SearchUserTopUps(userId int, keyword string, pageInfo *common.PageInfo) (to
if err = tx.Commit().Error; err != nil { if err = tx.Commit().Error; err != nil {
return nil, 0, err return nil, 0, err
} }
fillTopUpEmails(topups)
return topups, total, nil return topups, total, nil
} }


@@ -232,6 +262,7 @@ func SearchAllTopUps(keyword string, pageInfo *common.PageInfo) (topups []*TopUp
if err = tx.Commit().Error; err != nil { if err = tx.Commit().Error; err != nil {
return nil, 0, err return nil, 0, err
} }
fillTopUpEmails(topups)
return topups, total, nil return topups, total, nil
} }


@@ -269,7 +300,7 @@ func ManualCompleteTopUp(tradeNo string) error {
// 计算应充值额度: // 计算应充值额度:
// - Stripe/微信支付订单:Money 代表经分组倍率换算后的数量,直接 * QuotaPerUnit // - Stripe/微信支付订单:Money 代表经分组倍率换算后的数量,直接 * QuotaPerUnit
// - 其他订单(如易支付):Amount 为美元数量,* QuotaPerUnit // - 其他订单(如易支付):Amount 为美元数量,* QuotaPerUnit
if topUp.PaymentMethod == "stripe" || topUp.PaymentMethod == "wechat_pay" {
if topUp.PaymentMethod == "stripe" || topUp.PaymentMethod == "wechat_pay" || topUp.PaymentMethod == "alipay" {
dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit) dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit)
quotaToAdd = int(decimal.NewFromFloat(topUp.Money).Mul(dQuotaPerUnit).IntPart()) quotaToAdd = int(decimal.NewFromFloat(topUp.Money).Mul(dQuotaPerUnit).IntPart())
} else { } else {


+ 12
- 6
model/topup_wechat.go Переглянути файл

@@ -12,8 +12,18 @@ import (
) )


// RechargeWechat 微信支付充值完成(由回调触发) // RechargeWechat 微信支付充值完成(由回调触发)
// 与 Recharge/RechargeCreem 类似,使用事务+行锁保证幂等
func RechargeWechat(tradeNo string) error { func RechargeWechat(tradeNo string) error {
return rechargeByQRCodePayment(tradeNo, "微信支付")
}

// RechargeAlipay 支付宝充值完成(由回调触发)
func RechargeAlipay(tradeNo string) error {
return rechargeByQRCodePayment(tradeNo, "支付宝")
}

// rechargeByQRCodePayment 扫码支付充值完成(微信/支付宝通用)
// 使用事务+行锁保证幂等
func rechargeByQRCodePayment(tradeNo string, paymentMethod string) error {
if tradeNo == "" { if tradeNo == "" {
return errors.New("未提供支付单号") return errors.New("未提供支付单号")
} }
@@ -34,7 +44,6 @@ func RechargeWechat(tradeNo string) error {
} }


if topUp.Status == common.TopUpStatusSuccess { if topUp.Status == common.TopUpStatusSuccess {
// 已处理,幂等返回
return nil return nil
} }


@@ -48,9 +57,6 @@ func RechargeWechat(tradeNo string) error {
return err return err
} }


// 微信支付充值额度计算:
// topUp.Money = req.Amount * topUpGroupRatio(经分组倍率调整后的数量)
// 充值额度 = topUp.Money * QuotaPerUnit(与 Stripe 的 Recharge 逻辑一致)
dMoney := decimal.NewFromFloat(topUp.Money) dMoney := decimal.NewFromFloat(topUp.Money)
dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit) dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit)
quotaToAdd = dMoney.Mul(dQuotaPerUnit).IntPart() quotaToAdd = dMoney.Mul(dQuotaPerUnit).IntPart()
@@ -73,7 +79,7 @@ func RechargeWechat(tradeNo string) error {
} }


if quotaToAdd > 0 { if quotaToAdd > 0 {
RecordLog(userId, LogTypeTopup, fmt.Sprintf("使用微信支付充值成功,充值金额: %v,支付金额:%.2f", logger.FormatQuota(int(quotaToAdd)), payMoney))
RecordLog(userId, LogTypeTopup, fmt.Sprintf("使用%s充值成功,充值金额: %v,支付金额:%.2f", paymentMethod, logger.FormatQuota(int(quotaToAdd)), payMoney))
} }


return nil return nil


+ 40
- 6
model/user.go Переглянути файл

@@ -214,6 +214,20 @@ func GetMaxUserId() int {
return user.Id return user.Id
} }


// ApplySyncedQuota 将同步用户的 quota 替换为 synced_quota,
// 使 API 返回的 quota 对所有用户类型都表示实际可用余额。
func (u *User) ApplySyncedQuota() {
if u.IsSyncedUser() {
u.Quota = u.SyncedQuota
}
}

func applySyncedUserQuota(users []*User) {
for _, u := range users {
u.ApplySyncedQuota()
}
}

func GetAllUsers(pageInfo *common.PageInfo) (users []*User, total int64, err error) { func GetAllUsers(pageInfo *common.PageInfo) (users []*User, total int64, err error) {
// Start transaction // Start transaction
tx := DB.Begin() tx := DB.Begin()
@@ -245,6 +259,7 @@ func GetAllUsers(pageInfo *common.PageInfo) (users []*User, total int64, err err
return nil, 0, err return nil, 0, err
} }


applySyncedUserQuota(users)
return users, total, nil return users, total, nil
} }


@@ -312,6 +327,7 @@ func SearchUsers(keyword string, group string, startIdx int, num int) ([]*User,
return nil, 0, err return nil, 0, err
} }


applySyncedUserQuota(users)
return users, total, nil return users, total, nil
} }


@@ -410,7 +426,12 @@ func (user *User) Insert(inviterId int) error {
return err return err
} }
} }
user.Quota = common.QuotaForNewUser
matchedQuota := MatchEmailQuotaRule(user.Email)
if matchedQuota >= 0 {
user.Quota = int(matchedQuota)
} else {
user.Quota = common.QuotaForNewUser
}
//user.SetAccessToken(common.GetUUID()) //user.SetAccessToken(common.GetUUID())
user.AffCode = common.GetRandomString(4) user.AffCode = common.GetRandomString(4)


@@ -441,8 +462,12 @@ func (user *User) Insert(inviterId int) error {
} }
} }


if common.QuotaForNewUser > 0 {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", logger.LogQuota(common.QuotaForNewUser)))
if user.Quota > 0 {
if MatchEmailQuotaRule(user.Email) >= 0 {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s(邮箱后缀规则匹配)", logger.LogQuota(user.Quota)))
} else {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", logger.LogQuota(user.Quota)))
}
} }
if inviterId != 0 { if inviterId != 0 {
if common.QuotaForInvitee > 0 { if common.QuotaForInvitee > 0 {
@@ -469,7 +494,12 @@ func (user *User) InsertWithTx(tx *gorm.DB, inviterId int) error {
return err return err
} }
} }
user.Quota = common.QuotaForNewUser
matchedQuota := MatchEmailQuotaRule(user.Email)
if matchedQuota >= 0 {
user.Quota = int(matchedQuota)
} else {
user.Quota = common.QuotaForNewUser
}
user.AffCode = common.GetRandomString(4) user.AffCode = common.GetRandomString(4)


// 初始化用户设置 // 初始化用户设置
@@ -502,8 +532,12 @@ func (user *User) FinalizeOAuthUserCreation(inviterId int) {
} }
} }


if common.QuotaForNewUser > 0 {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", logger.LogQuota(common.QuotaForNewUser)))
if user.Quota > 0 {
if matchedQuota := MatchEmailQuotaRule(user.Email); matchedQuota >= 0 {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s(邮箱后缀规则匹配)", logger.LogQuota(user.Quota)))
} else {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", logger.LogQuota(user.Quota)))
}
} }
if inviterId != 0 { if inviterId != 0 {
if common.QuotaForInvitee > 0 { if common.QuotaForInvitee > 0 {


+ 18
- 2
model/vendor_meta.go Переглянути файл

@@ -18,6 +18,7 @@ type Vendor struct {
Description string `json:"description,omitempty" gorm:"type:text"` Description string `json:"description,omitempty" gorm:"type:text"`
Icon string `json:"icon,omitempty" gorm:"type:varchar(128)"` Icon string `json:"icon,omitempty" gorm:"type:varchar(128)"`
Status int `json:"status" gorm:"default:1"` Status int `json:"status" gorm:"default:1"`
SortOrder int `json:"sort_order" gorm:"default:999999"`
CreatedTime int64 `json:"created_time" gorm:"bigint"` CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"` UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
DeletedAt gorm.DeletedAt `json:"-" gorm:"index;uniqueIndex:uk_vendor_name_delete_at,priority:2"` DeletedAt gorm.DeletedAt `json:"-" gorm:"index;uniqueIndex:uk_vendor_name_delete_at,priority:2"`
@@ -65,7 +66,7 @@ func GetVendorByID(id int) (*Vendor, error) {
// GetAllVendors 获取全部供应商(分页) // GetAllVendors 获取全部供应商(分页)
func GetAllVendors(offset int, limit int) ([]*Vendor, error) { func GetAllVendors(offset int, limit int) ([]*Vendor, error) {
var vendors []*Vendor var vendors []*Vendor
err := DB.Offset(offset).Limit(limit).Find(&vendors).Error
err := DB.Order("sort_order ASC, id ASC").Offset(offset).Limit(limit).Find(&vendors).Error
return vendors, err return vendors, err
} }


@@ -81,8 +82,23 @@ func SearchVendors(keyword string, offset int, limit int) ([]*Vendor, int64, err
return nil, 0, err return nil, 0, err
} }
var vendors []*Vendor var vendors []*Vendor
if err := db.Offset(offset).Limit(limit).Order("id DESC").Find(&vendors).Error; err != nil {
if err := db.Order("sort_order ASC, id ASC").Offset(offset).Limit(limit).Find(&vendors).Error; err != nil {
return nil, 0, err return nil, 0, err
} }
return vendors, total, nil return vendors, total, nil
} }

// ReorderVendors 批量更新供应商排序值
func ReorderVendors(items []struct {
Id int `json:"id"`
SortOrder int `json:"sort_order"`
}) error {
return DB.Transaction(func(tx *gorm.DB) error {
for _, item := range items {
if err := tx.Model(&Vendor{}).Where("id = ?", item.Id).Update("sort_order", item.SortOrder).Error; err != nil {
return err
}
}
return nil
})
}

+ 2
- 2
relay/channel/api_request.go Переглянути файл

@@ -282,7 +282,7 @@ func DoApiRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBody
return nil, fmt.Errorf("get request url failed: %w", err) return nil, fmt.Errorf("get request url failed: %w", err)
} }
if common2.DebugEnabled { if common2.DebugEnabled {
println("fullRequestURL:", fullRequestURL)
common2.SysLog("fullRequestURL: " + fullRequestURL)
} }
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody) req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
if err != nil { if err != nil {
@@ -313,7 +313,7 @@ func DoFormRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBod
return nil, fmt.Errorf("get request url failed: %w", err) return nil, fmt.Errorf("get request url failed: %w", err)
} }
if common2.DebugEnabled { if common2.DebugEnabled {
println("fullRequestURL:", fullRequestURL)
common2.SysLog("fullRequestURL: " + fullRequestURL)
} }
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody) req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
if err != nil { if err != nil {


+ 4
- 4
relay/channel/claude/relay-claude.go Переглянути файл

@@ -635,7 +635,7 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d
if claudeResponse.Message != nil && claudeResponse.Message.Usage != nil { if claudeResponse.Message != nil && claudeResponse.Message.Usage != nil {
claudeInfo.Usage.PromptTokens = claudeResponse.Message.Usage.InputTokens claudeInfo.Usage.PromptTokens = claudeResponse.Message.Usage.InputTokens
claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Message.Usage.CacheReadInputTokens claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Message.Usage.CacheReadInputTokens
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Message.Usage.CacheCreationInputTokens
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Message.Usage.GetCacheCreationTotalTokens()
claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Message.Usage.GetCacheCreation5mTokens() claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Message.Usage.GetCacheCreation5mTokens()
claudeInfo.Usage.ClaudeCacheCreation1hTokens = claudeResponse.Message.Usage.GetCacheCreation1hTokens() claudeInfo.Usage.ClaudeCacheCreation1hTokens = claudeResponse.Message.Usage.GetCacheCreation1hTokens()
claudeInfo.Usage.CompletionTokens = claudeResponse.Message.Usage.OutputTokens claudeInfo.Usage.CompletionTokens = claudeResponse.Message.Usage.OutputTokens
@@ -659,8 +659,8 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d
if claudeResponse.Usage.CacheReadInputTokens > 0 { if claudeResponse.Usage.CacheReadInputTokens > 0 {
claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens
} }
if claudeResponse.Usage.CacheCreationInputTokens > 0 {
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Usage.CacheCreationInputTokens
if total := claudeResponse.Usage.GetCacheCreationTotalTokens(); total > 0 {
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = total
} }
if cacheCreation5m := claudeResponse.Usage.GetCacheCreation5mTokens(); cacheCreation5m > 0 { if cacheCreation5m := claudeResponse.Usage.GetCacheCreation5mTokens(); cacheCreation5m > 0 {
claudeInfo.Usage.ClaudeCacheCreation5mTokens = cacheCreation5m claudeInfo.Usage.ClaudeCacheCreation5mTokens = cacheCreation5m
@@ -811,7 +811,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
claudeInfo.Usage.CompletionTokens = claudeResponse.Usage.OutputTokens claudeInfo.Usage.CompletionTokens = claudeResponse.Usage.OutputTokens
claudeInfo.Usage.TotalTokens = claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens claudeInfo.Usage.TotalTokens = claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens
claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Usage.CacheCreationInputTokens
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Usage.GetCacheCreationTotalTokens()
claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Usage.GetCacheCreation5mTokens() claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Usage.GetCacheCreation5mTokens()
claudeInfo.Usage.ClaudeCacheCreation1hTokens = claudeResponse.Usage.GetCacheCreation1hTokens() claudeInfo.Usage.ClaudeCacheCreation1hTokens = claudeResponse.Usage.GetCacheCreation1hTokens()
} }


+ 27
- 11
relay/channel/codex/adaptor.go Переглянути файл

@@ -138,9 +138,21 @@ func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
if info.RelayMode != relayconstant.RelayModeResponses && info.RelayMode != relayconstant.RelayModeResponsesCompact { if info.RelayMode != relayconstant.RelayModeResponses && info.RelayMode != relayconstant.RelayModeResponsesCompact {
return "", errors.New("codex channel: only /v1/responses and /v1/responses/compact are supported") return "", errors.New("codex channel: only /v1/responses and /v1/responses/compact are supported")
} }
path := "/backend-api/codex/responses"

key := strings.TrimSpace(info.ApiKey)
if strings.HasPrefix(key, "{") {
// OAuth mode: route to ChatGPT backend API
path := "/backend-api/codex/responses"
if info.RelayMode == relayconstant.RelayModeResponsesCompact {
path = "/backend-api/codex/responses/compact"
}
return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, path, info.ChannelType), nil
}

// API key mode: route to standard /v1/responses
path := "/v1/responses"
if info.RelayMode == relayconstant.RelayModeResponsesCompact { if info.RelayMode == relayconstant.RelayModeResponsesCompact {
path = "/backend-api/codex/responses/compact"
path = "/v1/responses/compact"
} }
return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, path, info.ChannelType), nil return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, path, info.ChannelType), nil
} }
@@ -149,11 +161,20 @@ func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Header, info *rel
channel.SetupApiRequestHeader(info, c, req) channel.SetupApiRequestHeader(info, c, req)


key := strings.TrimSpace(info.ApiKey) key := strings.TrimSpace(info.ApiKey)
if !strings.HasPrefix(key, "{") {
return errors.New("codex channel: key must be a JSON object")
if strings.HasPrefix(key, "{") {
return setupOAuthHeader(req, key)
} }


oauthKey, err := ParseOAuthKey(key)
// Simple API key mode
req.Set("Authorization", "Bearer "+key)
if info.IsStream {
req.Set("Accept", "text/event-stream")
}
return nil
}

func setupOAuthHeader(req *http.Header, rawKey string) error {
oauthKey, err := ParseOAuthKey(rawKey)
if err != nil { if err != nil {
return err return err
} }
@@ -178,13 +199,8 @@ func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Header, info *rel
req.Set("originator", "codex_cli_rs") req.Set("originator", "codex_cli_rs")
} }


// chatgpt.com/backend-api/codex/responses is strict about Content-Type.
// Clients may omit it or include parameters like `application/json; charset=utf-8`,
// which can be rejected by the upstream. Force the exact media type.
req.Set("Content-Type", "application/json") req.Set("Content-Type", "application/json")
if info.IsStream {
req.Set("Accept", "text/event-stream")
} else if req.Get("Accept") == "" {
if req.Get("Accept") == "" {
req.Set("Accept", "application/json") req.Set("Accept", "application/json")
} }




+ 5
- 9
relay/channel/openai/chat_via_responses.go Переглянути файл

@@ -56,7 +56,9 @@ func OaiResponsesToChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
} }


if oaiError := responsesResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { if oaiError := responsesResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
apiErr := types.WithOpenAIError(*oaiError, resp.StatusCode)
apiErr.UpstreamBody = service.TruncateBody(string(body))
return nil, apiErr
} }


chatId := helper.GetResponseID(c) chatId := helper.GetResponseID(c)
@@ -484,14 +486,8 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
sentStop = true sentStop = true
} }


case "response.error", "response.failed":
if streamResp.Response != nil {
if oaiErr := streamResp.Response.GetOpenAIError(); oaiErr != nil && oaiErr.Type != "" {
streamErr = types.WithOpenAIError(*oaiErr, http.StatusInternalServerError)
return false
}
}
streamErr = types.NewOpenAIError(fmt.Errorf("responses stream error: %s", streamResp.Type), types.ErrorCodeBadResponse, http.StatusInternalServerError)
case "response.error", "response.failed", "error":
streamErr = handleResponsesStreamError(streamResp, data)
return false return false


default: default:


+ 11
- 4
relay/channel/openai/relay_responses.go Переглянути файл

@@ -31,7 +31,9 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
apiErr := types.WithOpenAIError(*oaiError, resp.StatusCode)
apiErr.UpstreamBody = service.TruncateBody(string(responseBody))
return nil, apiErr
} }


if responsesResponse.HasImageGenerationCall() { if responsesResponse.HasImageGenerationCall() {
@@ -78,10 +80,10 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp


var usage = &dto.Usage{} var usage = &dto.Usage{}
var responseTextBuilder strings.Builder var responseTextBuilder strings.Builder
var streamErr *types.NewAPIError


helper.StreamScannerHandler(c, resp, info, func(data string) bool { helper.StreamScannerHandler(c, resp, info, func(data string) bool {


// 检查当前数据是否包含 completed 状态和 usage 信息
var streamResponse dto.ResponsesStreamResponse var streamResponse dto.ResponsesStreamResponse
if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil { if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil {
sendResponsesStreamData(c, streamResponse, data) sendResponsesStreamData(c, streamResponse, data)
@@ -109,10 +111,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
} }
} }
case "response.output_text.delta": case "response.output_text.delta":
// 处理输出文本
responseTextBuilder.WriteString(streamResponse.Delta) responseTextBuilder.WriteString(streamResponse.Delta)
case dto.ResponsesOutputTypeItemDone: case dto.ResponsesOutputTypeItemDone:
// 函数调用处理
if streamResponse.Item != nil { if streamResponse.Item != nil {
switch streamResponse.Item.Type { switch streamResponse.Item.Type {
case dto.BuildInCallWebSearchCall: case dto.BuildInCallWebSearchCall:
@@ -123,6 +123,9 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
} }
} }
} }
case "response.error", "response.failed", "error":
streamErr = handleResponsesStreamError(streamResponse, data)
return false
} }
} else { } else {
logger.LogError(c, "failed to unmarshal stream response: "+err.Error()) logger.LogError(c, "failed to unmarshal stream response: "+err.Error())
@@ -130,6 +133,10 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
return true return true
}) })


if streamErr != nil {
return nil, streamErr
}

if usage.CompletionTokens == 0 { if usage.CompletionTokens == 0 {
// 计算输出文本的 token 数量 // 计算输出文本的 token 数量
tempStr := responseTextBuilder.String() tempStr := responseTextBuilder.String()


+ 29
- 0
relay/channel/openai/responses_error.go Переглянути файл

@@ -0,0 +1,29 @@
package openai

import (
"fmt"
"net/http"

"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/types"
)

// handleResponsesStreamError extracts error from a Responses API SSE event and attaches the raw data as UpstreamBody.
func handleResponsesStreamError(streamResp dto.ResponsesStreamResponse, data string) *types.NewAPIError {
if streamResp.Response != nil {
if oaiErr := streamResp.Response.GetOpenAIError(); oaiErr != nil && oaiErr.Type != "" {
err := types.WithOpenAIError(*oaiErr, http.StatusInternalServerError)
err.UpstreamBody = service.TruncateBody(data)
return err
}
}
if oaiErr := dto.GetOpenAIError(streamResp.Error); oaiErr != nil && oaiErr.Type != "" {
err := types.WithOpenAIError(*oaiErr, http.StatusInternalServerError)
err.UpstreamBody = service.TruncateBody(data)
return err
}
err := types.NewOpenAIError(fmt.Errorf("responses stream error: %s", streamResp.Type), types.ErrorCodeBadResponse, http.StatusInternalServerError)
err.UpstreamBody = service.TruncateBody(data)
return err
}

+ 64
- 0
relay/channel/openai/upstream_body_test.go Переглянути файл

@@ -0,0 +1,64 @@
package openai

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestResponsesStreamEventErrorParsing(t *testing.T) {
t.Run("standalone error event", func(t *testing.T) {
data := `{"type":"error","error":{"type":"too_many_requests","code":"too_many_requests","message":"Too Many Requests","param":null}}`

var streamResp dto.ResponsesStreamResponse
err := common.UnmarshalJsonStr(data, &streamResp)
require.NoError(t, err)
assert.Equal(t, "error", streamResp.Type)
assert.NotNil(t, streamResp.Error)

oaiErr := dto.GetOpenAIError(streamResp.Error)
require.NotNil(t, oaiErr)
assert.Equal(t, "too_many_requests", oaiErr.Type)
assert.Equal(t, "Too Many Requests", oaiErr.Message)
})

t.Run("response.failed event", func(t *testing.T) {
data := `{"type":"response.failed","response":{"id":"test-id","status":"failed","error":{"type":"server_error","message":"Internal server error"}}}`

var streamResp dto.ResponsesStreamResponse
err := common.UnmarshalJsonStr(data, &streamResp)
require.NoError(t, err)
assert.Equal(t, "response.failed", streamResp.Type)
require.NotNil(t, streamResp.Response)

oaiErr := streamResp.Response.GetOpenAIError()
require.NotNil(t, oaiErr)
assert.Equal(t, "server_error", oaiErr.Type)
})

t.Run("response.error event with error in response object", func(t *testing.T) {
data := `{"type":"response.error","response":{"id":"test-id","status":"failed","error":{"type":"invalid_request_error","message":"Invalid model"}}}`

var streamResp dto.ResponsesStreamResponse
err := common.UnmarshalJsonStr(data, &streamResp)
require.NoError(t, err)
assert.Equal(t, "response.error", streamResp.Type)

oaiErr := streamResp.Response.GetOpenAIError()
require.NotNil(t, oaiErr)
assert.Equal(t, "invalid_request_error", oaiErr.Type)
})

t.Run("normal event has no error", func(t *testing.T) {
data := `{"type":"response.output_text.delta","delta":"hello"}`

var streamResp dto.ResponsesStreamResponse
err := common.UnmarshalJsonStr(data, &streamResp)
require.NoError(t, err)
assert.Equal(t, "response.output_text.delta", streamResp.Type)
assert.Nil(t, streamResp.Error)
})
}

+ 1
- 1
relay/claude_handler.go Переглянути файл

@@ -160,7 +160,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} }


if common.DebugEnabled { if common.DebugEnabled {
println("requestBody: ", string(jsonData))
common.SysLog("claude requestBody: " + string(jsonData))
} }
requestBody = bytes.NewBuffer(jsonData) requestBody = bytes.NewBuffer(jsonData)
} }


+ 17
- 4
relay/compatible_handler.go Переглянути файл

@@ -27,6 +27,22 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )


func shouldUseChatCompletionsViaResponses(info *relaycommon.RelayInfo, passThroughGlobal bool) bool {
if info == nil {
return false
}
if info.RelayMode != relayconstant.RelayModeChatCompletions {
return false
}
if info.ChannelType == constant.ChannelTypeCodex {
return true
}
if passThroughGlobal || info.ChannelSetting.PassThroughBodyEnabled {
return false
}
return service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName)
}

func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)


@@ -76,10 +92,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
adaptor.Init(info) adaptor.Init(info)


passThroughGlobal := model_setting.GetGlobalSettings().PassThroughRequestEnabled passThroughGlobal := model_setting.GetGlobalSettings().PassThroughRequestEnabled
if info.RelayMode == relayconstant.RelayModeChatCompletions &&
!passThroughGlobal &&
!info.ChannelSetting.PassThroughBodyEnabled &&
service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) {
if shouldUseChatCompletionsViaResponses(info, passThroughGlobal) {
applySystemPromptIfNeeded(c, info, request) applySystemPromptIfNeeded(c, info, request)
usage, newApiErr := chatCompletionsViaResponses(c, info, adaptor, request) usage, newApiErr := chatCompletionsViaResponses(c, info, adaptor, request)
if newApiErr != nil { if newApiErr != nil {


+ 60
- 41
relay/helper/price.go Переглянути файл

@@ -4,6 +4,7 @@ import (
"fmt" "fmt"


"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger" "github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
@@ -14,9 +15,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )


// https://docs.claude.com/en/docs/build-with-claude/prompt-caching#1-hour-cache-duration
const claudeCacheCreation1hMultiplier = 6 / 3.75

// HandleGroupRatio checks for "auto_group" in the context and updates the group ratio and relayInfo.UsingGroup if present // HandleGroupRatio checks for "auto_group" in the context and updates the group ratio and relayInfo.UsingGroup if present
func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) types.GroupRatioInfo { func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) types.GroupRatioInfo {
groupRatioInfo := types.GroupRatioInfo{ groupRatioInfo := types.GroupRatioInfo{
@@ -52,17 +50,30 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
var modelRatio float64 var modelRatio float64
var completionRatio float64 var completionRatio float64
var channelPricingFound bool var channelPricingFound bool
var cacheRatio, cacheCreationRatio, imageRatio, audioRatio, audioCompletionRatio float64


// 尝试获取渠道定价(优先于全局定价) // 尝试获取渠道定价(优先于全局定价)
channelMetaAvailable := info != nil && info.ChannelMeta != nil && info.ChannelId > 0
if channelMetaAvailable {
cpRatio, cpCompletionRatio, cpPrice, cpUsePrice, found := model.GetEffectivePricing(info.OriginModelName, info.ChannelId)
if found {
modelRatio = cpRatio
completionRatio = cpCompletionRatio
modelPrice = cpPrice
usePrice = cpUsePrice
// ChannelMeta 在 InitChannelMeta 之前为 nil,但 Distribute 已将 channelId 写入 context
channelId := 0
if info != nil && info.ChannelMeta != nil && info.ChannelId > 0 {
channelId = info.ChannelId
} else {
channelId = common.GetContextKeyInt(c, constant.ContextKeyChannelId)
}
if channelId > 0 {
cp, found := model.GetEffectivePricing(info.OriginModelName, channelId)
if found && cp != nil {
modelRatio = cp.ModelRatio
completionRatio = cp.CompletionRatio
modelPrice = cp.ModelPrice
usePrice = cp.QuotaType == model.QuotaTypeByCall
channelPricingFound = true channelPricingFound = true
// 渠道定价的扩展比率(非零值直接使用)
cacheRatio = cp.CacheRatio
cacheCreationRatio = cp.CacheCreationRatio
imageRatio = cp.ImageRatio
audioRatio = cp.AudioRatio
audioCompletionRatio = cp.AudioCompletionRatio
} }
} }


@@ -74,13 +85,6 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
groupRatioInfo := HandleGroupRatio(c, info) groupRatioInfo := HandleGroupRatio(c, info)


var preConsumedQuota int var preConsumedQuota int
var cacheRatio float64
var imageRatio float64
var cacheCreationRatio float64
var cacheCreationRatio5m float64
var cacheCreationRatio1h float64
var audioRatio float64
var audioCompletionRatio float64
var freeModel bool var freeModel bool
if !usePrice { if !usePrice {
preConsumedTokens := common.Max(promptTokens, common.PreConsumedQuota) preConsumedTokens := common.Max(promptTokens, common.PreConsumedQuota)
@@ -103,14 +107,6 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
} }
completionRatio = ratio_setting.GetCompletionRatio(info.OriginModelName) completionRatio = ratio_setting.GetCompletionRatio(info.OriginModelName)
} }
cacheRatio, _ = ratio_setting.GetCacheRatio(info.OriginModelName)
cacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(info.OriginModelName)
cacheCreationRatio5m = cacheCreationRatio
// 固定1h和5min缓存写入价格的比例
cacheCreationRatio1h = cacheCreationRatio * claudeCacheCreation1hMultiplier
imageRatio, _ = ratio_setting.GetImageRatio(info.OriginModelName)
audioRatio = ratio_setting.GetAudioRatio(info.OriginModelName)
audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(info.OriginModelName)
ratio := modelRatio * groupRatioInfo.GroupRatio ratio := modelRatio * groupRatioInfo.GroupRatio
preConsumedQuota = int(float64(preConsumedTokens) * ratio) preConsumedQuota = int(float64(preConsumedTokens) * ratio)
} else { } else {
@@ -120,6 +116,24 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
preConsumedQuota = int(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio) preConsumedQuota = int(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio)
} }


// 全局比率作为回退(仅当渠道未设置、值为 0 时生效)
// 必须放在 usePrice 判断之外,因为 UpdatePriceDataForChannelPricing 可能改变 UsePrice
if cacheRatio == 0 {
cacheRatio, _ = ratio_setting.GetCacheRatio(info.OriginModelName)
}
if cacheCreationRatio == 0 {
cacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(info.OriginModelName)
}
if imageRatio == 0 {
imageRatio, _ = ratio_setting.GetImageRatio(info.OriginModelName)
}
if audioRatio == 0 {
audioRatio = ratio_setting.GetAudioRatio(info.OriginModelName)
}
if audioCompletionRatio == 0 {
audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(info.OriginModelName)
}

// check if free model pre-consume is disabled // check if free model pre-consume is disabled
if !operation_setting.GetQuotaSetting().EnableFreeModelPreConsume { if !operation_setting.GetQuotaSetting().EnableFreeModelPreConsume {
// if model price or ratio is 0, do not pre-consume quota // if model price or ratio is 0, do not pre-consume quota
@@ -146,14 +160,14 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
CompletionRatio: completionRatio, CompletionRatio: completionRatio,
GroupRatioInfo: groupRatioInfo, GroupRatioInfo: groupRatioInfo,
UsePrice: usePrice, UsePrice: usePrice,
QuotaToPreConsume: preConsumedQuota,
CacheRatio: cacheRatio, CacheRatio: cacheRatio,
ImageRatio: imageRatio, ImageRatio: imageRatio,
AudioRatio: audioRatio, AudioRatio: audioRatio,
AudioCompletionRatio: audioCompletionRatio, AudioCompletionRatio: audioCompletionRatio,
CacheCreationRatio: cacheCreationRatio, CacheCreationRatio: cacheCreationRatio,
CacheCreation5mRatio: cacheCreationRatio5m,
CacheCreation1hRatio: cacheCreationRatio1h,
QuotaToPreConsume: preConsumedQuota,
CacheCreation5mRatio: cacheCreationRatio,
CacheCreation1hRatio: cacheCreationRatio * types.ClaudeCacheCreation1hMultiplier,
} }


if common.DebugEnabled { if common.DebugEnabled {
@@ -217,30 +231,35 @@ func UpdatePriceDataForChannelPricing(c *gin.Context, info *relaycommon.RelayInf
return return
} }


cpRatio, cpCompletionRatio, cpPrice, cpUsePrice, found := model.GetEffectivePricing(info.OriginModelName, channelId)
cp, found := model.GetEffectivePricing(info.OriginModelName, channelId)
if !found { if !found {
return return
} }


// 更新 PriceData 中的定价相关字段
info.PriceData.ModelRatio = cpRatio
info.PriceData.CompletionRatio = cpCompletionRatio
info.PriceData.UsePrice = cpUsePrice
// 按次计费模式下使用渠道价格,按量计费模式下设置 ModelPrice = -1 让前端识别计费模式
if cpUsePrice {
info.PriceData.ModelPrice = cpPrice
info.PriceData.ModelRatio = cp.ModelRatio
info.PriceData.CompletionRatio = cp.CompletionRatio
info.PriceData.UsePrice = cp.QuotaType == model.QuotaTypeByCall
if info.PriceData.UsePrice {
info.PriceData.ModelPrice = cp.ModelPrice
} else { } else {
info.PriceData.ModelPrice = -1 info.PriceData.ModelPrice = -1
} }


// 重新计算预扣费额度(用于后续可能的引用)
if cpUsePrice {
info.PriceData.QuotaToPreConsume = int(cpPrice * common.QuotaPerUnit * info.PriceData.GroupRatioInfo.GroupRatio)
info.PriceData.ApplyChannelPricingRatios(cp.CacheRatio, cp.CacheCreationRatio, cp.ImageRatio, cp.AudioRatio, cp.AudioCompletionRatio)

if info.PriceData.UsePrice {
info.PriceData.QuotaToPreConsume = int(
cp.ModelPrice * common.QuotaPerUnit * info.PriceData.GroupRatioInfo.GroupRatio)
} else { } else {
estimateTokens := info.GetEstimatePromptTokens() estimateTokens := info.GetEstimatePromptTokens()
if estimateTokens > 0 { if estimateTokens > 0 {
ratio := cpRatio * info.PriceData.GroupRatioInfo.GroupRatio
ratio := cp.ModelRatio * info.PriceData.GroupRatioInfo.GroupRatio
info.PriceData.QuotaToPreConsume = int(float64(estimateTokens) * ratio) info.PriceData.QuotaToPreConsume = int(float64(estimateTokens) * ratio)
} }
} }

if common.DebugEnabled {
println(fmt.Sprintf("[ChannelPricing] updatePriceData: model=%s channel=%d modelRatio=%.4f completionRatio=%.4f cacheRatio=%.4f imageRatio=%.4f audioRatio=%.4f",
info.OriginModelName, channelId, cp.ModelRatio, cp.CompletionRatio, cp.CacheRatio, cp.ImageRatio, cp.AudioRatio))
}
} }

+ 22
- 0
router/api-router.go Переглянути файл

@@ -26,11 +26,14 @@ func SetApiRouter(router *gin.Engine) {
apiRouter.GET("/notice", controller.GetNotice) apiRouter.GET("/notice", controller.GetNotice)
apiRouter.GET("/user-agreement", controller.GetUserAgreement) apiRouter.GET("/user-agreement", controller.GetUserAgreement)
apiRouter.GET("/privacy-policy", controller.GetPrivacyPolicy) apiRouter.GET("/privacy-policy", controller.GetPrivacyPolicy)
apiRouter.GET("/terms", controller.GetTermsOfService)
apiRouter.GET("/usage-policy", controller.GetUsagePolicy)
apiRouter.GET("/about", controller.GetAbout) apiRouter.GET("/about", controller.GetAbout)
//apiRouter.GET("/midjourney", controller.GetMidjourney) //apiRouter.GET("/midjourney", controller.GetMidjourney)
apiRouter.GET("/home_page_content", controller.GetHomePageContent) apiRouter.GET("/home_page_content", controller.GetHomePageContent)
apiRouter.GET("/pricing", middleware.TryUserAuth(), controller.GetPricing) apiRouter.GET("/pricing", middleware.TryUserAuth(), controller.GetPricing)
apiRouter.GET("/channel-pricing/model/*name", middleware.TryUserAuth(), controller.GetChannelPricingByModelWithChannelInfo) apiRouter.GET("/channel-pricing/model/*name", middleware.TryUserAuth(), controller.GetChannelPricingByModelWithChannelInfo)
apiRouter.GET("/captcha", controller.GetCaptcha)
apiRouter.GET("/verification", middleware.EmailVerificationRateLimit(), middleware.TurnstileCheck(), controller.SendEmailVerification) apiRouter.GET("/verification", middleware.EmailVerificationRateLimit(), middleware.TurnstileCheck(), controller.SendEmailVerification)
apiRouter.GET("/reset_password", middleware.CriticalRateLimit(), middleware.TurnstileCheck(), controller.SendPasswordResetEmail) apiRouter.GET("/reset_password", middleware.CriticalRateLimit(), middleware.TurnstileCheck(), controller.SendPasswordResetEmail)
apiRouter.POST("/user/reset", middleware.CriticalRateLimit(), controller.ResetPassword) apiRouter.POST("/user/reset", middleware.CriticalRateLimit(), controller.ResetPassword)
@@ -49,6 +52,7 @@ func SetApiRouter(router *gin.Engine) {
apiRouter.POST("/stripe/webhook", controller.StripeWebhook) apiRouter.POST("/stripe/webhook", controller.StripeWebhook)
apiRouter.POST("/creem/webhook", controller.CreemWebhook) apiRouter.POST("/creem/webhook", controller.CreemWebhook)
apiRouter.POST("/wechat/pay/webhook", controller.WechatPayWebhook) apiRouter.POST("/wechat/pay/webhook", controller.WechatPayWebhook)
apiRouter.POST("/alipay/pay/webhook", controller.AlipayPayWebhook)
// Universal secure verification routes // Universal secure verification routes
apiRouter.POST("/verify", middleware.UserAuth(), middleware.CriticalRateLimit(), controller.UniversalVerify) apiRouter.POST("/verify", middleware.UserAuth(), middleware.CriticalRateLimit(), controller.UniversalVerify)


@@ -71,6 +75,7 @@ func SetApiRouter(router *gin.Engine) {
selfRoute.GET("/self/groups", controller.GetUserGroups) selfRoute.GET("/self/groups", controller.GetUserGroups)
selfRoute.GET("/self", controller.GetSelf) selfRoute.GET("/self", controller.GetSelf)
selfRoute.GET("/models", controller.GetUserModels) selfRoute.GET("/models", controller.GetUserModels)
selfRoute.GET("/model_channels", controller.GetModelChannels)
selfRoute.GET("/channels", controller.GetUserChannelsForBinding) selfRoute.GET("/channels", controller.GetUserChannelsForBinding)
selfRoute.PUT("/self", controller.UpdateSelf) selfRoute.PUT("/self", controller.UpdateSelf)
selfRoute.DELETE("/self", controller.DeleteSelf) selfRoute.DELETE("/self", controller.DeleteSelf)
@@ -93,6 +98,9 @@ func SetApiRouter(router *gin.Engine) {
selfRoute.POST("/wechat/pay/amount", controller.RequestWechatPayAmount) selfRoute.POST("/wechat/pay/amount", controller.RequestWechatPayAmount)
selfRoute.POST("/wechat/pay", controller.RequestWechatPay) selfRoute.POST("/wechat/pay", controller.RequestWechatPay)
selfRoute.GET("/wechat/pay/status", controller.WechatPayStatus) selfRoute.GET("/wechat/pay/status", controller.WechatPayStatus)
selfRoute.POST("/alipay/pay/amount", controller.RequestAlipayPayAmount)
selfRoute.POST("/alipay/pay", controller.RequestAlipayPay)
selfRoute.GET("/alipay/pay/status", controller.AlipayPayStatus)
selfRoute.POST("/aff_transfer", controller.TransferAffQuota) selfRoute.POST("/aff_transfer", controller.TransferAffQuota)
selfRoute.PUT("/setting", controller.UpdateUserSetting) selfRoute.PUT("/setting", controller.UpdateUserSetting)


@@ -190,6 +198,8 @@ func SetApiRouter(router *gin.Engine) {
channelPricingRoute.POST("/batch", controller.BatchCreateChannelPricing) channelPricingRoute.POST("/batch", controller.BatchCreateChannelPricing)
channelPricingRoute.POST("/copy_global/:channel_id", controller.CopyGlobalPricing) channelPricingRoute.POST("/copy_global/:channel_id", controller.CopyGlobalPricing)
channelPricingRoute.DELETE("/:id", controller.DeleteChannelPricing) channelPricingRoute.DELETE("/:id", controller.DeleteChannelPricing)
channelPricingRoute.POST("/set_default", controller.SetDefaultChannel)
channelPricingRoute.DELETE("/default/*name", controller.ClearDefaultChannel)
} }


// 定价标签路由(管理员权限) // 定价标签路由(管理员权限)
@@ -202,6 +212,16 @@ func SetApiRouter(router *gin.Engine) {
pricingTagRoute.DELETE("/:id", controller.DeletePricingTag) pricingTagRoute.DELETE("/:id", controller.DeletePricingTag)
} }


// 邮箱后缀额度规则路由(管理员权限)
emailQuotaRuleRoute := apiRouter.Group("/email_quota_rule")
emailQuotaRuleRoute.Use(middleware.AdminAuth())
{
emailQuotaRuleRoute.GET("/", controller.GetAllEmailQuotaRules)
emailQuotaRuleRoute.POST("/", controller.CreateEmailQuotaRule)
emailQuotaRuleRoute.PUT("/:id", controller.UpdateEmailQuotaRule)
emailQuotaRuleRoute.DELETE("/:id", controller.DeleteEmailQuotaRule)
}

// Custom OAuth provider management (root only) // Custom OAuth provider management (root only)
customOAuthRoute := apiRouter.Group("/custom-oauth-provider") customOAuthRoute := apiRouter.Group("/custom-oauth-provider")
customOAuthRoute.Use(middleware.RootAuth()) customOAuthRoute.Use(middleware.RootAuth())
@@ -350,6 +370,7 @@ func SetApiRouter(router *gin.Engine) {
vendorRoute.GET("/:id", controller.GetVendorMeta) vendorRoute.GET("/:id", controller.GetVendorMeta)
vendorRoute.POST("/", controller.CreateVendorMeta) vendorRoute.POST("/", controller.CreateVendorMeta)
vendorRoute.PUT("/", controller.UpdateVendorMeta) vendorRoute.PUT("/", controller.UpdateVendorMeta)
vendorRoute.PUT("/reorder", controller.ReorderVendors)
vendorRoute.DELETE("/:id", controller.DeleteVendorMeta) vendorRoute.DELETE("/:id", controller.DeleteVendorMeta)
} }


@@ -364,6 +385,7 @@ func SetApiRouter(router *gin.Engine) {
modelsRoute.GET("/:id", controller.GetModelMeta) modelsRoute.GET("/:id", controller.GetModelMeta)
modelsRoute.POST("/", controller.CreateModelMeta) modelsRoute.POST("/", controller.CreateModelMeta)
modelsRoute.PUT("/", controller.UpdateModelMeta) modelsRoute.PUT("/", controller.UpdateModelMeta)
modelsRoute.PUT("/reorder", controller.ReorderModels)
modelsRoute.DELETE("/:id", controller.DeleteModelMeta) modelsRoute.DELETE("/:id", controller.DeleteModelMeta)
} }




+ 8
- 0
router/web-router.go Переглянути файл

@@ -14,6 +14,14 @@ import (
) )


func SetWebRouter(router *gin.Engine, buildFS embed.FS, indexPage []byte) { func SetWebRouter(router *gin.Engine, buildFS embed.FS, indexPage []byte) {
// LOGO_FILE_PATH 优先级最高:注册 /logo.png 路由提供磁盘文件
// 必须在 static.Serve("/") 之前注册,否则会被 embed 静态文件拦截
if common.LogoFilePath != "" {
router.GET("/logo.png", func(c *gin.Context) {
c.File(common.LogoFilePath)
})
}

router.Use(gzip.Gzip(gzip.DefaultCompression)) router.Use(gzip.Gzip(gzip.DefaultCompression))
router.Use(middleware.GlobalWebRateLimit()) router.Use(middleware.GlobalWebRateLimit())
router.Use(middleware.Cache()) router.Use(middleware.Cache())


+ 6
- 7
service/channel_affinity_usage_cache_test.go Переглянути файл

@@ -4,7 +4,6 @@ import (
"fmt" "fmt"
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"time"


"github.com/QuantumNous/new-api/dto" "github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/types" "github.com/QuantumNous/new-api/types"
@@ -26,9 +25,9 @@ func buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP string)
} }


func TestObserveChannelAffinityUsageCacheByRelayFormat_ClaudeMode(t *testing.T) { func TestObserveChannelAffinityUsageCacheByRelayFormat_ClaudeMode(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default" usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP) ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)


usage := &dto.Usage{ usage := &dto.Usage{
@@ -53,9 +52,9 @@ func TestObserveChannelAffinityUsageCacheByRelayFormat_ClaudeMode(t *testing.T)
} }


func TestObserveChannelAffinityUsageCacheByRelayFormat_MixedMode(t *testing.T) { func TestObserveChannelAffinityUsageCacheByRelayFormat_MixedMode(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default" usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP) ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)


openAIUsage := &dto.Usage{ openAIUsage := &dto.Usage{
@@ -83,9 +82,9 @@ func TestObserveChannelAffinityUsageCacheByRelayFormat_MixedMode(t *testing.T) {
} }


func TestObserveChannelAffinityUsageCacheByRelayFormat_UnsupportedModeKeepsEmpty(t *testing.T) { func TestObserveChannelAffinityUsageCacheByRelayFormat_UnsupportedModeKeepsEmpty(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default" usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP) ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)


usage := &dto.Usage{ usage := &dto.Usage{


+ 2
- 24
service/codex_credential_refresh.go Переглянути файл

@@ -2,7 +2,6 @@ package service


import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"strings" "strings"
"time" "time"
@@ -16,28 +15,7 @@ type CodexCredentialRefreshOptions struct {
ResetCaches bool ResetCaches bool
} }


type CodexOAuthKey struct {
IDToken string `json:"id_token,omitempty"`
AccessToken string `json:"access_token,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`

AccountID string `json:"account_id,omitempty"`
LastRefresh string `json:"last_refresh,omitempty"`
Email string `json:"email,omitempty"`
Type string `json:"type,omitempty"`
Expired string `json:"expired,omitempty"`
}

func parseCodexOAuthKey(raw string) (*CodexOAuthKey, error) {
if strings.TrimSpace(raw) == "" {
return nil, errors.New("codex channel: empty oauth key")
}
var key CodexOAuthKey
if err := common.Unmarshal([]byte(raw), &key); err != nil {
return nil, errors.New("codex channel: invalid oauth key json")
}
return &key, nil
}
type CodexOAuthKey = common.CodexOAuthCredential


func RefreshCodexChannelCredential(ctx context.Context, channelID int, opts CodexCredentialRefreshOptions) (*CodexOAuthKey, *model.Channel, error) { func RefreshCodexChannelCredential(ctx context.Context, channelID int, opts CodexCredentialRefreshOptions) (*CodexOAuthKey, *model.Channel, error) {
ch, err := model.GetChannelById(channelID, true) ch, err := model.GetChannelById(channelID, true)
@@ -51,7 +29,7 @@ func RefreshCodexChannelCredential(ctx context.Context, channelID int, opts Code
return nil, nil, fmt.Errorf("channel type is not Codex") return nil, nil, fmt.Errorf("channel type is not Codex")
} }


oauthKey, err := parseCodexOAuthKey(strings.TrimSpace(ch.Key))
oauthKey, err := common.ParseCodexOAuthCredential(strings.TrimSpace(ch.Key))
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }


+ 1
- 1
service/codex_credential_refresh_task.go Переглянути файл

@@ -93,7 +93,7 @@ func runCodexCredentialAutoRefreshOnce() {
continue continue
} }


oauthKey, err := parseCodexOAuthKey(rawKey)
oauthKey, err := common.ParseCodexOAuthCredential(rawKey)
if err != nil { if err != nil {
continue continue
} }


+ 21
- 0
service/error.go Переглянути файл

@@ -58,6 +58,14 @@ func MidjourneyErrorWithStatusCodeWrapper(code int, desc string, statusCode int)
// return openaiErr // return openaiErr
//} //}


func TruncateBody(body string) string {
const maxBodyLen = 2048
if len(body) > maxBodyLen {
return body[:maxBodyLen] + "...(truncated)"
}
return body
}

func ClaudeErrorWrapper(err error, code string, statusCode int) *dto.ClaudeErrorWithStatusCode { func ClaudeErrorWrapper(err error, code string, statusCode int) *dto.ClaudeErrorWithStatusCode {
text := err.Error() text := err.Error()
lowerText := strings.ToLower(text) lowerText := strings.ToLower(text)
@@ -86,11 +94,20 @@ func ClaudeErrorWrapperLocal(err error, code string, statusCode int) *dto.Claude
func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFail bool) (newApiErr *types.NewAPIError) { func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFail bool) (newApiErr *types.NewAPIError) {
newApiErr = types.InitOpenAIError(types.ErrorCodeBadResponseStatusCode, resp.StatusCode) newApiErr = types.InitOpenAIError(types.ErrorCodeBadResponseStatusCode, resp.StatusCode)


// Capture upstream request-id from response headers
upstreamReqId := resp.Header.Get("request-id")
if upstreamReqId == "" {
upstreamReqId = resp.Header.Get("x-request-id")
}
newApiErr.UpstreamRequestId = upstreamReqId

responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return return
} }
CloseResponseBodyGracefully(resp) CloseResponseBodyGracefully(resp)
bodyStr := TruncateBody(string(responseBody))
newApiErr.UpstreamBody = bodyStr
var errResponse dto.GeneralErrorResponse var errResponse dto.GeneralErrorResponse
buildErrWithBody := func(message string) error { buildErrWithBody := func(message string) error {
if message == "" { if message == "" {
@@ -115,6 +132,8 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai
oaiError := errResponse.TryToOpenAIError() oaiError := errResponse.TryToOpenAIError()
if oaiError != nil { if oaiError != nil {
newApiErr = types.WithOpenAIError(*oaiError, resp.StatusCode) newApiErr = types.WithOpenAIError(*oaiError, resp.StatusCode)
newApiErr.UpstreamRequestId = upstreamReqId
newApiErr.UpstreamBody = bodyStr
if showBodyWhenFail { if showBodyWhenFail {
newApiErr.Err = buildErrWithBody(newApiErr.Error()) newApiErr.Err = buildErrWithBody(newApiErr.Error())
} }
@@ -122,6 +141,8 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai
} }
} }
newApiErr = types.NewOpenAIError(errors.New(errResponse.ToMessage()), types.ErrorCodeBadResponseStatusCode, resp.StatusCode) newApiErr = types.NewOpenAIError(errors.New(errResponse.ToMessage()), types.ErrorCodeBadResponseStatusCode, resp.StatusCode)
newApiErr.UpstreamRequestId = upstreamReqId
newApiErr.UpstreamBody = bodyStr
if showBodyWhenFail { if showBodyWhenFail {
newApiErr.Err = buildErrWithBody(newApiErr.Error()) newApiErr.Err = buildErrWithBody(newApiErr.Error())
} }


+ 32
- 0
service/truncate_body_test.go Переглянути файл

@@ -0,0 +1,32 @@
package service

import (
"strings"
"testing"

"github.com/stretchr/testify/assert"
)

func TestTruncateBody(t *testing.T) {
t.Parallel()

t.Run("short body unchanged", func(t *testing.T) {
body := `{"error":{"type":"too_many_requests","message":"Too Many Requests"}}`
assert.Equal(t, body, TruncateBody(body))
})

t.Run("empty body unchanged", func(t *testing.T) {
assert.Equal(t, "", TruncateBody(""))
})

t.Run("exact max length unchanged", func(t *testing.T) {
body := strings.Repeat("a", 2048)
assert.Equal(t, body, TruncateBody(body))
})

t.Run("over max length truncated", func(t *testing.T) {
body := strings.Repeat("a", 3000)
result := TruncateBody(body)
assert.Equal(t, strings.Repeat("a", 2048)+"...(truncated)", result)
})
}

+ 18
- 0
setting/payment_alipay.go Переглянути файл

@@ -0,0 +1,18 @@
package setting

var AlipayAppID = ""
var AlipayPrivateKey = "" // 应用私钥(RSA2)
var AlipayPublicKey = "" // 支付宝公钥(用于验签)
var AlipayNotifyURL = ""
var AlipayMinTopUp = 1
var AlipayUnitPrice = 7.0

// IsAlipayConfigured 检查支付宝核心配置是否完整
func IsAlipayConfigured() bool {
return AlipayAppID != "" &&
AlipayPrivateKey != "" &&
AlipayPublicKey != ""
}

// OnAlipayConfigChanged 配置变更时调用的回调函数(由 controller 包注册)
var OnAlipayConfigChanged func()

+ 18
- 0
setting/ratio_setting/model_ratio.go Переглянути файл

@@ -604,6 +604,15 @@ func GetAudioRatio(name string) float64 {
return 1 return 1
} }


func GetAudioRatioV2(name string) (float64, bool) {
name = FormatMatchingModelName(name)
ratio, ok := audioRatioMap.Get(name)
if !ok {
return 0, false
}
return ratio, true
}

func GetAudioCompletionRatio(name string) float64 { func GetAudioCompletionRatio(name string) float64 {
name = FormatMatchingModelName(name) name = FormatMatchingModelName(name)
if ratio, ok := audioCompletionRatioMap.Get(name); ok { if ratio, ok := audioCompletionRatioMap.Get(name); ok {
@@ -612,6 +621,15 @@ func GetAudioCompletionRatio(name string) float64 {
return 1 return 1
} }


func GetAudioCompletionRatioV2(name string) (float64, bool) {
name = FormatMatchingModelName(name)
ratio, ok := audioCompletionRatioMap.Get(name)
if !ok {
return 0, false
}
return ratio, true
}

func ContainsAudioRatio(name string) bool { func ContainsAudioRatio(name string) bool {
name = FormatMatchingModelName(name) name = FormatMatchingModelName(name)
_, ok := audioRatioMap.Get(name) _, ok := audioRatioMap.Get(name)


+ 16
- 4
setting/system_setting/legal.go Переглянути файл

@@ -3,13 +3,25 @@ package system_setting
import "github.com/QuantumNous/new-api/setting/config" import "github.com/QuantumNous/new-api/setting/config"


type LegalSettings struct { type LegalSettings struct {
UserAgreement string `json:"user_agreement"`
PrivacyPolicy string `json:"privacy_policy"`
UserAgreementZh string `json:"user_agreement_zh"`
UserAgreementEn string `json:"user_agreement_en"`
PrivacyPolicyZh string `json:"privacy_policy_zh"`
PrivacyPolicyEn string `json:"privacy_policy_en"`
TermsOfServiceZh string `json:"terms_of_service_zh"`
TermsOfServiceEn string `json:"terms_of_service_en"`
UsagePolicyZh string `json:"usage_policy_zh"`
UsagePolicyEn string `json:"usage_policy_en"`
} }


var defaultLegalSettings = LegalSettings{ var defaultLegalSettings = LegalSettings{
UserAgreement: "",
PrivacyPolicy: "",
UserAgreementZh: "",
UserAgreementEn: "",
PrivacyPolicyZh: "",
PrivacyPolicyEn: "",
TermsOfServiceZh: "",
TermsOfServiceEn: "",
UsagePolicyZh: "",
UsagePolicyEn: "",
} }


func init() { func init() {


+ 122
- 0
test-scripts/test_channel_pricing_extended.sh Переглянути файл

@@ -0,0 +1,122 @@
#!/bin/bash
set -e

BASE_URL="${1:-http://localhost:3000}"
ADMIN_KEY="${2}"

if [ -z "$ADMIN_KEY" ]; then
echo "Usage: $0 <base_url> <admin_key>"
exit 1
fi

echo "=== 渠道定价增强 - 集成测试 ==="

# ---------- 1. 创建渠道定价(含扩展字段) ----------
echo "--- Test 1: Create with extended ratios ---"
RESP=$(curl -s -X POST "$BASE_URL/api/channel-pricing/" \
-H "Authorization: Bearer $ADMIN_KEY" \
-H "Content-Type: application/json" \
-d '{
"model_name": "claude-3-5-sonnet",
"channel_id": 1,
"quota_type": 0,
"model_ratio": 3.0,
"completion_ratio": 15.0,
"cache_ratio": 0.5,
"cache_creation_ratio": 0.625,
"image_ratio": 1.5,
"audio_ratio": 2.0,
"audio_completion_ratio": 1.8
}')
echo "$RESP" | python3 -m json.tool

CACHE_RATIO=$(echo "$RESP" | python3 -c "import sys,json; print(json.load(sys.stdin)['data']['cache_ratio'])")
if [ "$CACHE_RATIO" = "0.5" ]; then
echo " [PASS] cache_ratio = 0.5"
else
echo " [FAIL] cache_ratio expected 0.5, got $CACHE_RATIO"
fi

# ---------- 2. 查询渠道定价(验证新字段返回) ----------
echo "--- Test 2: Query and verify extended fields ---"
RESP=$(curl -s "$BASE_URL/api/channel-pricing/model/claude-3-5-sonnet" \
-H "Authorization: Bearer $ADMIN_KEY")
echo "$RESP" | python3 -m json.tool

# ---------- 3. 更新渠道定价(修改扩展字段) ----------
echo "--- Test 3: Update extended ratios ---"
RESP=$(curl -s -X POST "$BASE_URL/api/channel-pricing/" \
-H "Authorization: Bearer $ADMIN_KEY" \
-H "Content-Type: application/json" \
-d '{
"model_name": "claude-3-5-sonnet",
"channel_id": 1,
"quota_type": 0,
"model_ratio": 3.0,
"completion_ratio": 15.0,
"cache_ratio": 0.8,
"cache_creation_ratio": 0.0
}')
echo "$RESP" | python3 -m json.tool

# ---------- 4. 验证回退:未设置的字段为 0 ----------
echo "--- Test 4: Verify unset fields = 0 ---"
CACHE_CREATION=$(echo "$RESP" | python3 -c "import sys,json; print(json.load(sys.stdin)['data']['cache_creation_ratio'])")
if [ "$CACHE_CREATION" = "0.0" ]; then
echo " [PASS] cache_creation_ratio = 0 (unset)"
else
echo " [FAIL] cache_creation_ratio expected 0.0, got $CACHE_CREATION"
fi

# ---------- 5. 负值校验 ----------
echo "--- Test 5: Reject negative values ---"
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" -X POST "$BASE_URL/api/channel-pricing/" \
-H "Authorization: Bearer $ADMIN_KEY" \
-H "Content-Type: application/json" \
-d '{
"model_name": "claude-3-5-sonnet",
"channel_id": 1,
"quota_type": 0,
"model_ratio": 3.0,
"cache_ratio": -1.0
}')
if [ "$HTTP_CODE" = "400" ] || [ "$HTTP_CODE" = "422" ]; then
echo " [PASS] negative value rejected with HTTP $HTTP_CODE"
else
echo " [FAIL] expected 400/422, got HTTP $HTTP_CODE"
fi

# ---------- 6. 用户端查询(带渠道信息 + CASE WHEN 回退) ----------
echo "--- Test 6: User-facing query with global fallback ---"
RESP=$(curl -s "$BASE_URL/api/channel-pricing/model/claude-3-5-sonnet" \
-H "Authorization: Bearer $ADMIN_KEY")
echo "$RESP" | python3 -c "
import sys, json
data = json.load(sys.stdin)
for item in data.get('data', []):
ch = item.get('channel_id', '?')
cr = item.get('cache_ratio', 'N/A')
ccr = item.get('cache_creation_ratio', 'N/A')
ir = item.get('image_ratio', 'N/A')
print(f' channel={ch} cache_ratio={cr} cache_creation_ratio={ccr} image_ratio={ir}')
"

# ---------- 7. 删除 + 验证缓存清除 ----------
echo "--- Test 7: Delete and verify cache cleared ---"
ID=$(curl -s "$BASE_URL/api/channel-pricing/model/claude-3-5-sonnet" \
-H "Authorization: Bearer $ADMIN_KEY" | \
python3 -c "import sys,json; items=json.load(sys.stdin).get('data',[]); print(items[0]['id'] if items else '')")
if [ -n "$ID" ]; then
curl -s -X DELETE "$BASE_URL/api/channel-pricing/$ID" \
-H "Authorization: Bearer $ADMIN_KEY"
echo " Deleted pricing id=$ID"

RESP=$(curl -s "$BASE_URL/api/channel-pricing/model/claude-3-5-sonnet" \
-H "Authorization: Bearer $ADMIN_KEY")
COUNT=$(echo "$RESP" | python3 -c "import sys,json; print(len(json.load(sys.stdin).get('data',[])))")
echo " Remaining pricings for claude-3-5-sonnet: $COUNT"
else
echo " [SKIP] No pricing to delete"
fi

echo "=== 集成测试完成 ==="

+ 10
- 8
types/error.go Переглянути файл

@@ -88,14 +88,16 @@ const (
) )


type NewAPIError struct { type NewAPIError struct {
Err error
RelayError any
skipRetry bool
recordErrorLog *bool
errorType ErrorType
errorCode ErrorCode
StatusCode int
Metadata json.RawMessage
Err error
RelayError any
skipRetry bool
recordErrorLog *bool
errorType ErrorType
errorCode ErrorCode
StatusCode int
Metadata json.RawMessage
UpstreamRequestId string
UpstreamBody string
} }


// Unwrap enables errors.Is / errors.As to work with NewAPIError by exposing the underlying error. // Unwrap enables errors.Is / errors.As to work with NewAPIError by exposing the underlying error.


+ 26
- 0
types/price_data.go Переглянути файл

@@ -2,6 +2,10 @@ package types


import "fmt" import "fmt"


// ClaudeCacheCreation1hMultiplier 1小时缓存写入价格相对于5分钟的比例
// https://docs.claude.com/en/docs/build-with-claude/prompt-caching#1-hour-cache-duration
const ClaudeCacheCreation1hMultiplier = 6 / 3.75

type GroupRatioInfo struct { type GroupRatioInfo struct {
GroupRatio float64 GroupRatio float64
GroupSpecialRatio float64 GroupSpecialRatio float64
@@ -27,6 +31,28 @@ type PriceData struct {
GroupRatioInfo GroupRatioInfo GroupRatioInfo GroupRatioInfo
} }


// ApplyChannelPricingRatios 将渠道定价的扩展比率应用到 PriceData(非零值覆盖)
// cacheRatio, cacheCreationRatio, imageRatio, audioRatio, audioCompletionRatio
func (p *PriceData) ApplyChannelPricingRatios(cacheRatio, cacheCreationRatio, imageRatio, audioRatio, audioCompletionRatio float64) {
if cacheRatio != 0 {
p.CacheRatio = cacheRatio
}
if cacheCreationRatio != 0 {
p.CacheCreationRatio = cacheCreationRatio
p.CacheCreation5mRatio = cacheCreationRatio
p.CacheCreation1hRatio = cacheCreationRatio * ClaudeCacheCreation1hMultiplier
}
if imageRatio != 0 {
p.ImageRatio = imageRatio
}
if audioRatio != 0 {
p.AudioRatio = audioRatio
}
if audioCompletionRatio != 0 {
p.AudioCompletionRatio = audioCompletionRatio
}
}

func (p *PriceData) AddOtherRatio(key string, ratio float64) { func (p *PriceData) AddOtherRatio(key string, ratio float64) {
if p.OtherRatios == nil { if p.OtherRatios == nil {
p.OtherRatios = make(map[string]float64) p.OtherRatios = make(map[string]float64)


+ 1
- 1
web/index.html Переглянути файл

@@ -10,7 +10,7 @@
content="OpenAI 接口聚合管理,支持多种渠道包括 Azure,可用于二次分发管理 key,仅单可执行文件,已打包好 Docker 镜像,一键部署,开箱即用" content="OpenAI 接口聚合管理,支持多种渠道包括 Azure,可用于二次分发管理 key,仅单可执行文件,已打包好 Docker 镜像,一键部署,开箱即用"
/> />
<meta name="generator" content="new-api" /> <meta name="generator" content="new-api" />
<title>New API</title>
<title>Loading...</title>
<!--umami--> <!--umami-->
<!--Google Analytics--> <!--Google Analytics-->
</head> </head>


+ 18
- 0
web/src/App.jsx Переглянути файл

@@ -55,6 +55,8 @@ const Dashboard = lazy(() => import('./pages/Dashboard'));
const About = lazy(() => import('./pages/About')); const About = lazy(() => import('./pages/About'));
const UserAgreement = lazy(() => import('./pages/UserAgreement')); const UserAgreement = lazy(() => import('./pages/UserAgreement'));
const PrivacyPolicy = lazy(() => import('./pages/PrivacyPolicy')); const PrivacyPolicy = lazy(() => import('./pages/PrivacyPolicy'));
const Terms = lazy(() => import('./pages/Terms'));
const UsagePolicy = lazy(() => import('./pages/UsagePolicy'));


function DynamicOAuth2Callback() { function DynamicOAuth2Callback() {
const { provider } = useParams(); const { provider } = useParams();
@@ -358,6 +360,22 @@ function App() {
</Suspense> </Suspense>
} }
/> />
<Route
path='/user-agreement'
element={
<Suspense fallback={<Loading></Loading>} key={location.pathname}>
<Terms />
</Suspense>
}
/>
<Route
path='/privacy-policy'
element={
<Suspense fallback={<Loading></Loading>} key={location.pathname}>
<UsagePolicy />
</Suspense>
}
/>
<Route <Route
path='/console/chat/:id?' path='/console/chat/:id?'
element={ element={


+ 21
- 11
web/src/components/auth/LoginForm.jsx Переглянути файл

@@ -17,8 +17,8 @@ along with this program. If not, see <https://www.gnu.org/licenses/>.
For commercial licensing, please contact support@quantumnous.com For commercial licensing, please contact support@quantumnous.com
*/ */


import React, { useContext, useEffect, useMemo, useRef, useState } from 'react';
import { Link, useNavigate, useSearchParams } from 'react-router-dom';
import React, { useCallback, useContext, useEffect, useMemo, useRef, useState } from 'react';
import { Link, useLocation, useNavigate, useSearchParams } from 'react-router-dom';
import { UserContext } from '../../context/User'; import { UserContext } from '../../context/User';
import { StatusContext } from '../../context/Status'; import { StatusContext } from '../../context/Status';
import { import {
@@ -69,6 +69,7 @@ import { SiDiscord } from 'react-icons/si';


const LoginForm = () => { const LoginForm = () => {
let navigate = useNavigate(); let navigate = useNavigate();
const location = useLocation();
const { t } = useTranslation(); const { t } = useTranslation();
const githubButtonTextKeyByState = { const githubButtonTextKeyByState = {
idle: '使用 GitHub 继续', idle: '使用 GitHub 继续',
@@ -113,6 +114,15 @@ const LoginForm = () => {
const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]); const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]);
const [customOAuthLoading, setCustomOAuthLoading] = useState({}); const [customOAuthLoading, setCustomOAuthLoading] = useState({});


const navigateAfterLogin = useCallback(() => {
const from = location.state?.from;
if (from && from.pathname) {
navigate(from.pathname + (from.search || ''));
} else {
navigate('/console');
}
}, [location.state, navigate]);

const logo = getLogo(); const logo = getLogo();
const systemName = getSystemName(); const systemName = getSystemName();


@@ -198,7 +208,7 @@ const LoginForm = () => {
localStorage.setItem('user', JSON.stringify(data)); localStorage.setItem('user', JSON.stringify(data));
setUserData(data); setUserData(data);
updateAPI(); updateAPI();
navigate('/');
navigateAfterLogin();
showSuccess(t('登录成功!')); showSuccess(t('登录成功!'));
setShowWeChatLoginModal(false); setShowWeChatLoginModal(false);
} else { } else {
@@ -255,7 +265,7 @@ const LoginForm = () => {
centered: true, centered: true,
}); });
} }
navigate('/console');
navigateAfterLogin();
} else { } else {
showError(message); showError(message);
} }
@@ -300,7 +310,7 @@ const LoginForm = () => {
showSuccess(t('登录成功!')); showSuccess(t('登录成功!'));
setUserData(data); setUserData(data);
updateAPI(); updateAPI();
navigate('/');
navigateAfterLogin();
} else { } else {
showError(message); showError(message);
} }
@@ -456,7 +466,7 @@ const LoginForm = () => {
setUserData(finish.data); setUserData(finish.data);
updateAPI(); updateAPI();
showSuccess(t('登录成功!')); showSuccess(t('登录成功!'));
navigate('/console');
navigateAfterLogin();
} else { } else {
showError(finish.message || t('Passkey 登录失败,请重试')); showError(finish.message || t('Passkey 登录失败,请重试'));
} }
@@ -490,8 +500,8 @@ const LoginForm = () => {
userDispatch({ type: 'login', payload: data }); userDispatch({ type: 'login', payload: data });
setUserData(data); setUserData(data);
updateAPI(); updateAPI();
showSuccess('登录成功!');
navigate('/console');
showSuccess(t('登录成功!'));
navigateAfterLogin();
}; };


// 返回登录页面 // 返回登录页面
@@ -505,7 +515,7 @@ const LoginForm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3} className='!text-gray-800'> <Title heading={3} className='!text-gray-800'>
{systemName} {systemName}
</Title> </Title>
@@ -721,7 +731,7 @@ const LoginForm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3}>{systemName}</Title> <Title heading={3}>{systemName}</Title>
</div> </div>


@@ -885,7 +895,7 @@ const LoginForm = () => {
}} }}
> >
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<img src={status.wechat_qrcode} alt={t('微信二维码')} className='mb-4' />
<img src={status.wechat_qrcode} alt={t('微信二维码')} referrerPolicy='no-referrer' crossOrigin='anonymous' className='mb-4' />
</div> </div>


<div className='text-center mb-4'> <div className='text-center mb-4'>


+ 1
- 1
web/src/components/auth/PasswordResetConfirm.jsx Переглянути файл

@@ -118,7 +118,7 @@ const PasswordResetConfirm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3} className='!text-gray-800'> <Title heading={3} className='!text-gray-800'>
{systemName} {systemName}
</Title> </Title>


+ 1
- 1
web/src/components/auth/PasswordResetForm.jsx Переглянути файл

@@ -118,7 +118,7 @@ const PasswordResetForm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3} className='!text-gray-800'> <Title heading={3} className='!text-gray-800'>
{systemName} {systemName}
</Title> </Title>


+ 68
- 9
web/src/components/auth/RegisterForm.jsx Переглянути файл

@@ -108,6 +108,9 @@ const RegisterForm = () => {
const [hasPrivacyPolicy, setHasPrivacyPolicy] = useState(false); const [hasPrivacyPolicy, setHasPrivacyPolicy] = useState(false);
const [githubButtonState, setGithubButtonState] = useState('idle'); const [githubButtonState, setGithubButtonState] = useState('idle');
const [githubButtonDisabled, setGithubButtonDisabled] = useState(false); const [githubButtonDisabled, setGithubButtonDisabled] = useState(false);
const [captchaId, setCaptchaId] = useState('');
const [captchaImage, setCaptchaImage] = useState('');
const [captchaCode, setCaptchaCode] = useState('');
const githubTimeoutRef = useRef(null); const githubTimeoutRef = useRef(null);
const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]); const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]);


@@ -141,6 +144,8 @@ const RegisterForm = () => {
hasCustomOAuthProviders, hasCustomOAuthProviders,
); );


const captchaEnabled = !!status?.captcha_enabled;

const [showEmailVerification, setShowEmailVerification] = useState(false); const [showEmailVerification, setShowEmailVerification] = useState(false);


useEffect(() => { useEffect(() => {
@@ -176,6 +181,26 @@ const RegisterForm = () => {
}; };
}, []); }, []);


const loadCaptcha = async () => {
try {
const res = await API.get('/api/captcha');
const { success, data } = res.data;
if (success) {
setCaptchaId(data.id);
setCaptchaImage(data.captcha_image);
setCaptchaCode('');
}
} catch (error) {
// silent fail
}
};

useEffect(() => {
if (showEmailVerification && captchaEnabled) {
loadCaptcha();
}
}, [showEmailVerification, captchaEnabled]);

const onWeChatLoginClicked = () => { const onWeChatLoginClicked = () => {
setWechatLoading(true); setWechatLoading(true);
setShowWeChatLoginModal(true); setShowWeChatLoginModal(true);
@@ -184,7 +209,7 @@ const RegisterForm = () => {


const onSubmitWeChatVerificationCode = async () => { const onSubmitWeChatVerificationCode = async () => {
if (turnstileEnabled && turnstileToken === '') { if (turnstileEnabled && turnstileToken === '') {
showInfo('请稍后几秒重试,Turnstile 正在检查用户环境!');
showInfo(t('请稍后几秒重试,Turnstile 正在检查用户环境!'));
return; return;
} }
setWechatCodeSubmitLoading(true); setWechatCodeSubmitLoading(true);
@@ -257,21 +282,28 @@ const RegisterForm = () => {
const sendVerificationCode = async () => { const sendVerificationCode = async () => {
if (inputs.email === '') return; if (inputs.email === '') return;
if (turnstileEnabled && turnstileToken === '') { if (turnstileEnabled && turnstileToken === '') {
showInfo('请稍后几秒重试,Turnstile 正在检查用户环境!');
showInfo(t('请稍后几秒重试,Turnstile 正在检查用户环境!'));
return;
}
if (captchaEnabled && captchaCode === '') {
showInfo(t('请先输入图片验证码'));
return; return;
} }
setVerificationCodeLoading(true); setVerificationCodeLoading(true);
try { try {
const res = await API.get(
`/api/verification?email=${encodeURIComponent(inputs.email)}&turnstile=${turnstileToken}`,
);
let url = `/api/verification?email=${encodeURIComponent(inputs.email)}&turnstile=${turnstileToken}`;
if (captchaEnabled) {
url += `&captcha_id=${encodeURIComponent(captchaId)}&captcha_code=${encodeURIComponent(captchaCode)}`;
}
const res = await API.get(url);
const { success, message } = res.data; const { success, message } = res.data;
if (success) { if (success) {
showSuccess(t('验证码发送成功,请检查你的邮箱!')); showSuccess(t('验证码发送成功,请检查你的邮箱!'));
setDisableButton(true); // 发送成功后禁用按钮,开始倒计时
setDisableButton(true);
} else { } else {
showError(message); showError(message);
} }
if (captchaEnabled) loadCaptcha();
} catch (error) { } catch (error) {
showError(t('发送验证码失败,请重试')); showError(t('发送验证码失败,请重试'));
} finally { } finally {
@@ -396,7 +428,7 @@ const RegisterForm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3} className='!text-gray-800'> <Title heading={3} className='!text-gray-800'>
{systemName} {systemName}
</Title> </Title>
@@ -559,7 +591,7 @@ const RegisterForm = () => {
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<div className='w-full max-w-md'> <div className='w-full max-w-md'>
<div className='flex items-center justify-center mb-6 gap-2'> <div className='flex items-center justify-center mb-6 gap-2'>
<img src={logo} alt='Logo' className='h-10 rounded-full' />
<img src={logo} alt='Logo' referrerPolicy='no-referrer' crossOrigin='anonymous' className='h-10 rounded-full' />
<Title heading={3} className='!text-gray-800'> <Title heading={3} className='!text-gray-800'>
{systemName} {systemName}
</Title> </Title>
@@ -624,6 +656,33 @@ const RegisterForm = () => {
</Button> </Button>
} }
/> />
{captchaEnabled && (
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<Form.Input
field='captcha_code'
label={t('图片验证码')}
placeholder={t('请输入图片验证码')}
name='captcha_code'
style={{ flex: 1 }}
onChange={(value) => setCaptchaCode(value)}
value={captchaCode}
prefix={<IconKey />}
/>
<img
src={captchaImage}
alt='captcha'
onClick={loadCaptcha}
style={{
height: 40,
cursor: 'pointer',
borderRadius: 4,
border: '1px solid #e0e0e0',
marginTop: 22,
}}
title={t('点击刷新验证码')}
/>
</div>
)}
<Form.Input <Form.Input
field='verification_code' field='verification_code'
label={t('验证码')} label={t('验证码')}
@@ -745,7 +804,7 @@ const RegisterForm = () => {
}} }}
> >
<div className='flex flex-col items-center'> <div className='flex flex-col items-center'>
<img src={status.wechat_qrcode} alt='微信二维码' className='mb-4' />
<img src={status.wechat_qrcode} alt={t('微信二维码')} referrerPolicy='no-referrer' crossOrigin='anonymous' className='mb-4' />
</div> </div>


<div className='text-center mb-4'> <div className='text-center mb-4'>


+ 3
- 2
web/src/components/common/markdown/MarkdownRenderer.jsx Переглянути файл

@@ -268,7 +268,7 @@ export function PreCode(props) {
color: 'var(--semi-color-text-2)', color: 'var(--semi-color-text-2)',
}} }}
> >
HTML预览:
{t('HTML预览:')}
</div> </div>
<SandboxedHtmlPreview code={htmlCode} /> <SandboxedHtmlPreview code={htmlCode} />
</div> </div>
@@ -635,6 +635,7 @@ function _MarkdownContent(props) {
export const MarkdownContent = React.memo(_MarkdownContent); export const MarkdownContent = React.memo(_MarkdownContent);


export function MarkdownRenderer(props) { export function MarkdownRenderer(props) {
const { t } = useTranslation();
const { const {
content, content,
loading, loading,
@@ -680,7 +681,7 @@ export function MarkdownRenderer(props) {
animation: 'spin 1s linear infinite', animation: 'spin 1s linear infinite',
}} }}
/> />
正在渲染...
{t('正在渲染...')}
</div> </div>
) : ( ) : (
<MarkdownContent <MarkdownContent


+ 1
- 1
web/src/components/common/ui/JSONEditor.jsx Переглянути файл

@@ -661,7 +661,7 @@ const JSONEditor = ({
{hasJsonError && ( {hasJsonError && (
<Banner <Banner
type='danger' type='danger'
description={`JSON 格式错误: ${jsonError}`}
description={`${t('JSON 格式错误')}: ${jsonError}`}
className='mb-3' className='mb-3'
/> />
)} )}


+ 2
- 0
web/src/components/layout/Footer.jsx Переглянути файл

@@ -52,6 +52,8 @@ const FooterBar = () => {
<img <img
src={logo} src={logo}
alt={systemName} alt={systemName}
referrerPolicy='no-referrer'
crossOrigin='anonymous'
className='w-16 h-16 rounded-full bg-gray-800 p-1.5 object-contain' className='w-16 h-16 rounded-full bg-gray-800 p-1.5 object-contain'
/> />
</div> </div>


+ 5
- 4
web/src/components/layout/PageLayout.jsx Переглянути файл

@@ -91,6 +91,11 @@ const PageLayout = () => {
if (success) { if (success) {
statusDispatch({ type: 'set', payload: data }); statusDispatch({ type: 'set', payload: data });
setStatusData(data); setStatusData(data);
// Apply admin-configured default language only if user has no preference
const savedLang = localStorage.getItem('i18nextLng');
if (data.default_language && !savedLang) {
i18n.changeLanguage(data.default_language);
}
} else { } else {
showError('Unable to connect to server'); showError('Unable to connect to server');
} }
@@ -113,10 +118,6 @@ const PageLayout = () => {
linkElement.href = logo; linkElement.href = logo;
} }
} }
const savedLang = localStorage.getItem('i18nextLng');
if (savedLang) {
i18n.changeLanguage(savedLang);
}
}, [i18n]); }, [i18n]);


return ( return (


+ 2
- 0
web/src/components/layout/headerbar/HeaderLogo.jsx Переглянути файл

@@ -44,6 +44,8 @@ const HeaderLogo = ({
<img <img
src={logo} src={logo}
alt='logo' alt='logo'
referrerPolicy='no-referrer'
crossOrigin='anonymous'
className={`absolute inset-0 w-full h-full transition-all duration-200 group-hover:scale-110 rounded-full ${!isLoading && logoLoaded ? 'opacity-100' : 'opacity-0'}`} className={`absolute inset-0 w-full h-full transition-all duration-200 group-hover:scale-110 rounded-full ${!isLoading && logoLoaded ? 'opacity-100' : 'opacity-0'}`}
/> />
</div> </div>


+ 0
- 30
web/src/components/layout/headerbar/LanguageSelector.jsx Переглянути файл

@@ -27,7 +27,6 @@ const LanguageSelector = ({ currentLang, onLanguageChange, t }) => {
position='bottomRight' position='bottomRight'
render={ render={
<Dropdown.Menu className='!bg-semi-color-bg-overlay !border-semi-color-border !shadow-lg !rounded-lg dark:!bg-gray-700 dark:!border-gray-600'> <Dropdown.Menu className='!bg-semi-color-bg-overlay !border-semi-color-border !shadow-lg !rounded-lg dark:!bg-gray-700 dark:!border-gray-600'>
{/* Language sorting: Order by English name (Chinese, English, French, Japanese, Russian) */}
<Dropdown.Item <Dropdown.Item
onClick={() => onLanguageChange('zh-CN')} onClick={() => onLanguageChange('zh-CN')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'zh-CN' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`} className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'zh-CN' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
@@ -35,40 +34,11 @@ const LanguageSelector = ({ currentLang, onLanguageChange, t }) => {
简体中文 简体中文
</Dropdown.Item> </Dropdown.Item>
<Dropdown.Item <Dropdown.Item
onClick={() => onLanguageChange('zh-TW')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'zh-TW' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
>
繁體中文
</Dropdown.Item> <Dropdown.Item
onClick={() => onLanguageChange('en')} onClick={() => onLanguageChange('en')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'en' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`} className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'en' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
> >
English English
</Dropdown.Item> </Dropdown.Item>
<Dropdown.Item
onClick={() => onLanguageChange('fr')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'fr' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
>
Français
</Dropdown.Item>
<Dropdown.Item
onClick={() => onLanguageChange('ja')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'ja' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
>
日本語
</Dropdown.Item>
<Dropdown.Item
onClick={() => onLanguageChange('ru')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'ru' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
>
Русский
</Dropdown.Item>
<Dropdown.Item
onClick={() => onLanguageChange('vi')}
className={`!px-3 !py-1.5 !text-sm !text-semi-color-text-0 dark:!text-gray-200 ${currentLang === 'vi' ? '!bg-semi-color-primary-light-default dark:!bg-blue-600 !font-semibold' : 'hover:!bg-semi-color-fill-1 dark:hover:!bg-gray-600'}`}
>
Tiếng Việt
</Dropdown.Item>
</Dropdown.Menu> </Dropdown.Menu>
} }
> >


+ 1
- 1
web/src/components/playground/CodeViewer.jsx Переглянути файл

@@ -201,7 +201,7 @@ const CodeViewer = ({ content, title, language = 'json' }) => {
} }
return ( return (
formattedContent.substring(0, PERFORMANCE_CONFIG.PREVIEW_LENGTH) + formattedContent.substring(0, PERFORMANCE_CONFIG.PREVIEW_LENGTH) +
'\n\n// ... 内容被截断以提升性能 ...'
'\n\n// ... ' + t('内容被截断以提升性能') + ' ...'
); );
}, [formattedContent, contentMetrics.isLarge, isExpanded]); }, [formattedContent, contentMetrics.isLarge, isExpanded]);




+ 1
- 1
web/src/components/playground/DebugPanel.jsx Переглянути файл

@@ -146,7 +146,7 @@ const DebugPanel = ({
{t('预览请求体')} {t('预览请求体')}
{customRequestMode && ( {customRequestMode && (
<span className='px-1.5 py-0.5 text-xs bg-orange-100 text-orange-600 rounded-full'> <span className='px-1.5 py-0.5 text-xs bg-orange-100 text-orange-600 rounded-full'>
自定义
{t('自定义')}
</span> </span>
)} )}
</div> </div>


+ 2
- 2
web/src/components/playground/MessageContent.jsx Переглянути файл

@@ -272,7 +272,7 @@ const MessageContent = ({
<div key={index} className='max-w-sm'> <div key={index} className='max-w-sm'>
<img <img
src={imgItem.image_url.url} src={imgItem.image_url.url}
alt={`用户上传的图片 ${index + 1}`}
alt={t('用户上传的图片', { index: index + 1 })}
className='rounded-lg max-w-full h-auto shadow-sm border' className='rounded-lg max-w-full h-auto shadow-sm border'
style={{ maxHeight: '300px' }} style={{ maxHeight: '300px' }}
onError={(e) => { onError={(e) => {
@@ -284,7 +284,7 @@ const MessageContent = ({
className='text-red-500 text-sm p-2 bg-red-50 rounded-lg border border-red-200' className='text-red-500 text-sm p-2 bg-red-50 rounded-lg border border-red-200'
style={{ display: 'none' }} style={{ display: 'none' }}
> >
图片加载失败: {imgItem.image_url.url}
{t('图片加载失败')}: {imgItem.image_url.url}
</div> </div>
</div> </div>
))} ))}


+ 2
- 1
web/src/components/playground/OptimizedComponents.js Переглянути файл

@@ -74,7 +74,8 @@ export const OptimizedSettingsPanel = React.memo(
prevProps.showSettings === nextProps.showSettings && prevProps.showSettings === nextProps.showSettings &&
JSON.stringify(prevProps.previewPayload) === JSON.stringify(prevProps.previewPayload) ===
JSON.stringify(nextProps.previewPayload) && JSON.stringify(nextProps.previewPayload) &&
JSON.stringify(prevProps.messages) === JSON.stringify(nextProps.messages)
JSON.stringify(prevProps.messages) === JSON.stringify(nextProps.messages) &&
JSON.stringify(prevProps.channels) === JSON.stringify(nextProps.channels)
); );
}, },
); );


+ 33
- 0
web/src/components/playground/SettingsPanel.jsx Переглянути файл

@@ -45,6 +45,7 @@ const SettingsPanel = ({
onCustomRequestBodyChange, onCustomRequestBodyChange,
previewPayload, previewPayload,
messages, messages,
channels = [],
}) => { }) => {
const { t } = useTranslation(); const { t } = useTranslation();


@@ -176,6 +177,38 @@ const SettingsPanel = ({
/> />
</div> </div>


{/* 通道选择 */}
<div className={customRequestMode ? 'opacity-50' : ''}>
<div className='flex items-center gap-2 mb-2'>
<Typography.Text strong className='text-sm'>
{t('通道')}
</Typography.Text>
{customRequestMode && (
<Typography.Text className='text-xs text-orange-600'>
({t('已在自定义模式中忽略')})
</Typography.Text>
)}
</div>
<Select
placeholder={t('请选择通道')}
name='channelId'
selection
filter={selectFilter}
autoClearSearchValue={false}
onChange={(value) => onInputChange('channelId', value)}
value={inputs.channelId}
autoComplete='new-password'
optionList={channels.map((ch) => ({
value: ch.id,
label: ch.public_name || ch.name,
}))}
style={{ width: '100%' }}
dropdownStyle={{ width: '100%', maxWidth: '100%' }}
className='!rounded-lg'
disabled={customRequestMode || channels.length === 0}
/>
</div>

{/* 图片URL输入 */} {/* 图片URL输入 */}
<div className={customRequestMode ? 'opacity-50' : ''}> <div className={customRequestMode ? 'opacity-50' : ''}>
<ImageUrlInput <ImageUrlInput


+ 2
- 2
web/src/components/playground/ThinkingContent.jsx Переглянути файл

@@ -105,7 +105,7 @@ const ThinkingContent = ({
style={{ color: 'white' }} style={{ color: 'white' }}
className='text-xs mt-0.5 opacity-80 hidden sm:block' className='text-xs mt-0.5 opacity-80 hidden sm:block'
> >
来源: {thinkingSource}
{t('来源')}: {thinkingSource}
</Typography.Text> </Typography.Text>
)} )}
</div> </div>
@@ -122,7 +122,7 @@ const ThinkingContent = ({
style={{ color: 'white' }} style={{ color: 'white' }}
className='text-xs sm:text-sm font-medium opacity-90' className='text-xs sm:text-sm font-medium opacity-90'
> >
思考中
{t('思考中')}
</Typography.Text> </Typography.Text>
</div> </div>
)} )}


+ 5
- 4
web/src/components/playground/configStorage.js Переглянути файл

@@ -21,6 +21,7 @@ import {
STORAGE_KEYS, STORAGE_KEYS,
DEFAULT_CONFIG, DEFAULT_CONFIG,
} from '../../constants/playground.constants'; } from '../../constants/playground.constants';
import i18next from 'i18next';


const MESSAGES_STORAGE_KEY = 'playground_messages'; const MESSAGES_STORAGE_KEY = 'playground_messages';


@@ -215,16 +216,16 @@ export const importConfig = (file) => {


resolve(importedConfig); resolve(importedConfig);
} else { } else {
reject(new Error('配置文件格式无效'));
reject(new Error(i18next.t('配置文件格式无效')));
} }
} catch (parseError) { } catch (parseError) {
reject(new Error('解析配置文件失败: ' + parseError.message));
reject(new Error(i18next.t('解析配置文件失败: ') + parseError.message));
} }
}; };
reader.onerror = () => reject(new Error('读取文件失败'));
reader.onerror = () => reject(new Error(i18next.t('读取文件失败')));
reader.readAsText(file); reader.readAsText(file);
} catch (error) { } catch (error) {
reject(new Error('导入配置失败: ' + error.message));
reject(new Error(i18next.t('导入配置失败: ') + error.message));
} }
}); });
}; };

+ 1
- 1
web/src/components/settings/ModelDeploymentSetting.jsx Переглянути файл

@@ -60,7 +60,7 @@ const ModelDeploymentSetting = () => {
setLoading(true); setLoading(true);
await getOptions(); await getOptions();
} catch (error) { } catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
console.error(error); console.error(error);
} finally { } finally {
setLoading(false); setLoading(false);


+ 1
- 1
web/src/components/settings/ModelSetting.jsx Переглянути файл

@@ -95,7 +95,7 @@ const ModelSetting = () => {
await getOptions(); await getOptions();
// showSuccess('刷新成功'); // showSuccess('刷新成功');
} catch (error) { } catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
console.error(error); console.error(error);
} finally { } finally {
setLoading(false); setLoading(false);


+ 137
- 53
web/src/components/settings/OtherSetting.jsx Переглянути файл

@@ -21,6 +21,7 @@ import React, { useContext, useEffect, useRef, useState } from 'react';
import { import {
Banner, Banner,
Button, Button,
ButtonGroup,
Col, Col,
Form, Form,
Row, Row,
@@ -34,15 +35,26 @@ import { useTranslation } from 'react-i18next';
import { StatusContext } from '../../context/Status'; import { StatusContext } from '../../context/Status';
import Text from '@douyinfe/semi-ui/lib/es/typography/text'; import Text from '@douyinfe/semi-ui/lib/es/typography/text';


const LEGAL_USER_AGREEMENT_KEY = 'legal.user_agreement';
const LEGAL_PRIVACY_POLICY_KEY = 'legal.privacy_policy';
const LEGAL_KEYS = {
userAgreement: { zh: 'legal.user_agreement_zh', en: 'legal.user_agreement_en' },
privacyPolicy: { zh: 'legal.privacy_policy_zh', en: 'legal.privacy_policy_en' },
termsOfService: { zh: 'legal.terms_of_service_zh', en: 'legal.terms_of_service_en' },
usagePolicy: { zh: 'legal.usage_policy_zh', en: 'legal.usage_policy_en' },
};


const OtherSetting = () => { const OtherSetting = () => {
const { t } = useTranslation(); const { t } = useTranslation();
const [editingLang, setEditingLang] = useState('zh');
let [inputs, setInputs] = useState({ let [inputs, setInputs] = useState({
Notice: '', Notice: '',
[LEGAL_USER_AGREEMENT_KEY]: '',
[LEGAL_PRIVACY_POLICY_KEY]: '',
[LEGAL_KEYS.userAgreement.zh]: '',
[LEGAL_KEYS.userAgreement.en]: '',
[LEGAL_KEYS.privacyPolicy.zh]: '',
[LEGAL_KEYS.privacyPolicy.en]: '',
[LEGAL_KEYS.termsOfService.zh]: '',
[LEGAL_KEYS.termsOfService.en]: '',
[LEGAL_KEYS.usagePolicy.zh]: '',
[LEGAL_KEYS.usagePolicy.en]: '',
SystemName: '', SystemName: '',
Logo: '', Logo: '',
Footer: '', Footer: '',
@@ -74,8 +86,10 @@ const OtherSetting = () => {


const [loadingInput, setLoadingInput] = useState({ const [loadingInput, setLoadingInput] = useState({
Notice: false, Notice: false,
[LEGAL_USER_AGREEMENT_KEY]: false,
[LEGAL_PRIVACY_POLICY_KEY]: false,
userAgreement: false,
privacyPolicy: false,
termsOfService: false,
usagePolicy: false,
SystemName: false, SystemName: false,
Logo: false, Logo: false,
HomePageContent: false, HomePageContent: false,
@@ -88,6 +102,13 @@ const OtherSetting = () => {
setInputs((inputs) => ({ ...inputs, [name]: value })); setInputs((inputs) => ({ ...inputs, [name]: value }));
}; };


// 语言切换时同步 form values,确保重新挂载的 TextArea 拿到正确内容
useEffect(() => {
if (formAPISettingGeneral.current) {
formAPISettingGeneral.current.setValues(inputs);
}
}, [editingLang]);

// 通用设置 // 通用设置
const formAPISettingGeneral = useRef(); const formAPISettingGeneral = useRef();
// 通用设置 - Notice // 通用设置 - Notice
@@ -103,48 +124,19 @@ const OtherSetting = () => {
setLoadingInput((loadingInput) => ({ ...loadingInput, Notice: false })); setLoadingInput((loadingInput) => ({ ...loadingInput, Notice: false }));
} }
}; };
// 通用设置 - UserAgreement
const submitUserAgreement = async () => {
// 通用法律文档保存(同时保存中英文)
const submitLegalDoc = async (docKey, successMsg, errorMsg) => {
const keys = LEGAL_KEYS[docKey];
try { try {
setLoadingInput((loadingInput) => ({
...loadingInput,
[LEGAL_USER_AGREEMENT_KEY]: true,
}));
await updateOption(
LEGAL_USER_AGREEMENT_KEY,
inputs[LEGAL_USER_AGREEMENT_KEY],
);
showSuccess(t('用户协议已更新'));
setLoadingInput((prev) => ({ ...prev, [docKey]: true }));
await updateOption(keys.zh, inputs[keys.zh]);
await updateOption(keys.en, inputs[keys.en]);
showSuccess(t(successMsg));
} catch (error) { } catch (error) {
console.error(t('用户协议更新失败'), error);
showError(t('用户协议更新失败'));
console.error(t(errorMsg), error);
showError(t(errorMsg));
} finally { } finally {
setLoadingInput((loadingInput) => ({
...loadingInput,
[LEGAL_USER_AGREEMENT_KEY]: false,
}));
}
};
// 通用设置 - PrivacyPolicy
const submitPrivacyPolicy = async () => {
try {
setLoadingInput((loadingInput) => ({
...loadingInput,
[LEGAL_PRIVACY_POLICY_KEY]: true,
}));
await updateOption(
LEGAL_PRIVACY_POLICY_KEY,
inputs[LEGAL_PRIVACY_POLICY_KEY],
);
showSuccess(t('隐私政策已更新'));
} catch (error) {
console.error(t('隐私政策更新失败'), error);
showError(t('隐私政策更新失败'));
} finally {
setLoadingInput((loadingInput) => ({
...loadingInput,
[LEGAL_PRIVACY_POLICY_KEY]: false,
}));
setLoadingInput((prev) => ({ ...prev, [docKey]: false }));
} }
}; };
// 个性化设置 // 个性化设置
@@ -376,11 +368,26 @@ const OtherSetting = () => {
{t('设置公告')} {t('设置公告')}
</Button> </Button>
<Form.TextArea <Form.TextArea
label={t('用户协议')}
label={
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<span>{t('用户协议')}</span>
<ButtonGroup size='small'>
<Button
type={editingLang === 'zh' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('zh')}
>中文</Button>
<Button
type={editingLang === 'en' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('en')}
>English</Button>
</ButtonGroup>
</div>
}
placeholder={t( placeholder={t(
'在此输入用户协议内容,支持 Markdown & HTML 代码', '在此输入用户协议内容,支持 Markdown & HTML 代码',
)} )}
field={LEGAL_USER_AGREEMENT_KEY}
field={LEGAL_KEYS.userAgreement[editingLang]}
key={`ua_${editingLang}`}
onChange={handleInputChange} onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }} style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }} autosize={{ minRows: 6, maxRows: 12 }}
@@ -389,17 +396,32 @@ const OtherSetting = () => {
)} )}
/> />
<Button <Button
onClick={submitUserAgreement}
loading={loadingInput[LEGAL_USER_AGREEMENT_KEY]}
onClick={() => submitLegalDoc('userAgreement', '用户协议已更新', '用户协议更新失败')}
loading={loadingInput['userAgreement']}
> >
{t('设置用户协议')} {t('设置用户协议')}
</Button> </Button>
<Form.TextArea <Form.TextArea
label={t('隐私政策')}
label={
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<span>{t('隐私政策')}</span>
<ButtonGroup size='small'>
<Button
type={editingLang === 'zh' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('zh')}
>中文</Button>
<Button
type={editingLang === 'en' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('en')}
>English</Button>
</ButtonGroup>
</div>
}
placeholder={t( placeholder={t(
'在此输入隐私政策内容,支持 Markdown & HTML 代码', '在此输入隐私政策内容,支持 Markdown & HTML 代码',
)} )}
field={LEGAL_PRIVACY_POLICY_KEY}
field={LEGAL_KEYS.privacyPolicy[editingLang]}
key={`pp_${editingLang}`}
onChange={handleInputChange} onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }} style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }} autosize={{ minRows: 6, maxRows: 12 }}
@@ -408,11 +430,73 @@ const OtherSetting = () => {
)} )}
/> />
<Button <Button
onClick={submitPrivacyPolicy}
loading={loadingInput[LEGAL_PRIVACY_POLICY_KEY]}
onClick={() => submitLegalDoc('privacyPolicy', '隐私政策已更新', '隐私政策更新失败')}
loading={loadingInput['privacyPolicy']}
> >
{t('设置隐私政策')} {t('设置隐私政策')}
</Button> </Button>
<Form.TextArea
label={
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<span>{t('服务条款')}</span>
<ButtonGroup size='small'>
<Button
type={editingLang === 'zh' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('zh')}
>中文</Button>
<Button
type={editingLang === 'en' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('en')}
>English</Button>
</ButtonGroup>
</div>
}
placeholder={t(
'在此输入服务条款内容,支持 Markdown & HTML 代码',
)}
field={LEGAL_KEYS.termsOfService[editingLang]}
key={`tos_${editingLang}`}
onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }}
/>
<Button
onClick={() => submitLegalDoc('termsOfService', '服务条款已更新', '服务条款更新失败')}
loading={loadingInput['termsOfService']}
>
{t('设置服务条款')}
</Button>
<Form.TextArea
label={
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<span>{t('使用政策')}</span>
<ButtonGroup size='small'>
<Button
type={editingLang === 'zh' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('zh')}
>中文</Button>
<Button
type={editingLang === 'en' ? 'primary' : 'tertiary'}
onClick={() => setEditingLang('en')}
>English</Button>
</ButtonGroup>
</div>
}
placeholder={t(
'在此输入使用政策内容,支持 Markdown & HTML 代码',
)}
field={LEGAL_KEYS.usagePolicy[editingLang]}
key={`up_${editingLang}`}
onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }}
/>
<Button
onClick={() => submitLegalDoc('usagePolicy', '使用政策已更新', '使用政策更新失败')}
loading={loadingInput['usagePolicy']}
>
{t('设置使用政策')}
</Button>
</Form.Section> </Form.Section>
</Card> </Card>
</Form> </Form>


+ 6
- 0
web/src/components/settings/PaymentSetting.jsx Переглянути файл

@@ -24,6 +24,7 @@ import SettingsPaymentGateway from '../../pages/Setting/Payment/SettingsPaymentG
import SettingsPaymentGatewayStripe from '../../pages/Setting/Payment/SettingsPaymentGatewayStripe'; import SettingsPaymentGatewayStripe from '../../pages/Setting/Payment/SettingsPaymentGatewayStripe';
import SettingsPaymentGatewayCreem from '../../pages/Setting/Payment/SettingsPaymentGatewayCreem'; import SettingsPaymentGatewayCreem from '../../pages/Setting/Payment/SettingsPaymentGatewayCreem';
import SettingsPaymentGatewayWechat from '../../pages/Setting/Payment/SettingsPaymentGatewayWechat'; import SettingsPaymentGatewayWechat from '../../pages/Setting/Payment/SettingsPaymentGatewayWechat';
import SettingsPaymentGatewayAlipay from '../../pages/Setting/Payment/SettingsPaymentGatewayAlipay';
import { API, showError, toBoolean } from '../../helpers'; import { API, showError, toBoolean } from '../../helpers';
import { useTranslation } from 'react-i18next'; import { useTranslation } from 'react-i18next';


@@ -101,6 +102,8 @@ const PaymentSetting = () => {
case 'StripeMinTopUp': case 'StripeMinTopUp':
case 'WechatPayUnitPrice': case 'WechatPayUnitPrice':
case 'WechatPayMinTopUp': case 'WechatPayMinTopUp':
case 'AlipayUnitPrice':
case 'AlipayMinTopUp':
newInputs[item.key] = parseFloat(item.value); newInputs[item.key] = parseFloat(item.value);
break; break;
default: default:
@@ -152,6 +155,9 @@ const PaymentSetting = () => {
<Card style={{ marginTop: '10px' }}> <Card style={{ marginTop: '10px' }}>
<SettingsPaymentGatewayWechat options={inputs} refresh={onRefresh} /> <SettingsPaymentGatewayWechat options={inputs} refresh={onRefresh} />
</Card> </Card>
<Card style={{ marginTop: '10px' }}>
<SettingsPaymentGatewayAlipay options={inputs} refresh={onRefresh} />
</Card>
</Spin> </Spin>
</> </>
); );


+ 1
- 1
web/src/components/settings/RateLimitSetting.jsx Переглянути файл

@@ -64,7 +64,7 @@ const RateLimitSetting = () => {
await getOptions(); await getOptions();
// showSuccess('刷新成功'); // showSuccess('刷新成功');
} catch (error) { } catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
} finally { } finally {
setLoading(false); setLoading(false);
} }


+ 1
- 1
web/src/components/settings/RatioSetting.jsx Переглянути файл

@@ -83,7 +83,7 @@ const RatioSetting = () => {
setLoading(true); setLoading(true);
await getOptions(); await getOptions();
} catch (error) { } catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
} finally { } finally {
setLoading(false); setLoading(false);
} }


+ 40
- 1
web/src/components/settings/SystemSetting.jsx Переглянути файл

@@ -78,6 +78,7 @@ const SystemSetting = () => {
WeChatServerToken: '', WeChatServerToken: '',
WeChatAccountQRCodeImageURL: '', WeChatAccountQRCodeImageURL: '',
TurnstileCheckEnabled: '', TurnstileCheckEnabled: '',
CaptchaEnabled: '',
TurnstileSiteKey: '', TurnstileSiteKey: '',
TurnstileSecretKey: '', TurnstileSecretKey: '',
RegisterEnabled: '', RegisterEnabled: '',
@@ -100,6 +101,7 @@ const SystemSetting = () => {
LinuxDOClientSecret: '', LinuxDOClientSecret: '',
LinuxDOMinimumTrustLevel: '', LinuxDOMinimumTrustLevel: '',
ServerAddress: '', ServerAddress: '',
DefaultLanguage: '',
// SSRF防护配置 // SSRF防护配置
'fetch_setting.enable_ssrf_protection': true, 'fetch_setting.enable_ssrf_protection': true,
'fetch_setting.allow_private_ip': '', 'fetch_setting.allow_private_ip': '',
@@ -179,6 +181,7 @@ const SystemSetting = () => {
case 'TelegramOAuthEnabled': case 'TelegramOAuthEnabled':
case 'RegisterEnabled': case 'RegisterEnabled':
case 'TurnstileCheckEnabled': case 'TurnstileCheckEnabled':
case 'CaptchaEnabled':
case 'EmailDomainRestrictionEnabled': case 'EmailDomainRestrictionEnabled':
case 'EmailAliasRestrictionEnabled': case 'EmailAliasRestrictionEnabled':
case 'SMTPSSLEnabled': case 'SMTPSSLEnabled':
@@ -317,6 +320,10 @@ const SystemSetting = () => {
await updateOptions([{ key: 'ServerAddress', value: ServerAddress }]); await updateOptions([{ key: 'ServerAddress', value: ServerAddress }]);
}; };


const submitDefaultLanguage = async () => {
await updateOptions([{ key: 'DefaultLanguage', value: inputs.DefaultLanguage || '' }]);
};

const submitSMTP = async () => { const submitSMTP = async () => {
const options = []; const options = [];


@@ -716,7 +723,7 @@ const SystemSetting = () => {
<Row <Row
gutter={{ xs: 8, sm: 16, md: 24, lg: 24, xl: 24, xxl: 24 }} gutter={{ xs: 8, sm: 16, md: 24, lg: 24, xl: 24, xxl: 24 }}
> >
<Col xs={24} sm={24} md={24} lg={24} xl={24}>
<Col xs={24} sm={24} md={24} lg={12} xl={12}>
<Form.Input <Form.Input
field='ServerAddress' field='ServerAddress'
label={t('服务器地址')} label={t('服务器地址')}
@@ -726,10 +733,33 @@ const SystemSetting = () => {
)} )}
/> />
</Col> </Col>
<Col xs={24} sm={24} md={24} lg={12} xl={12}>
<Form.Select
field='DefaultLanguage'
label={t('默认语言')}
placeholder={t('未设置时跟随浏览器语言')}
optionList={[
{ label: t('自动(跟随浏览器)'), value: '' },
{ label: '简体中文', value: 'zh-CN' },
{ label: '繁體中文', value: 'zh-TW' },
{ label: 'English', value: 'en' },
{ label: 'Français', value: 'fr' },
{ label: '日本語', value: 'ja' },
{ label: 'Русский', value: 'ru' },
{ label: 'Tiếng Việt', value: 'vi' },
]}
extraText={t(
'设置后,未登录用户和未设置语言偏好的已登录用户将强制使用此语言',
)}
/>
</Col>
</Row> </Row>
<Button onClick={submitServerAddress}> <Button onClick={submitServerAddress}>
{t('更新服务器地址')} {t('更新服务器地址')}
</Button> </Button>
<Button onClick={submitDefaultLanguage}>
{t('保存默认语言')}
</Button>
</Form.Section> </Form.Section>
</Card> </Card>


@@ -1033,6 +1063,15 @@ const SystemSetting = () => {
> >
{t('允许 Turnstile 用户校验')} {t('允许 Turnstile 用户校验')}
</Form.Checkbox> </Form.Checkbox>
<Form.Checkbox
field='CaptchaEnabled'
noLabel
onChange={(e) =>
handleCheckboxChange('CaptchaEnabled', e)
}
>
{t('图片验证码')}
</Form.Checkbox>
</Col> </Col>
<Col xs={24} sm={24} md={12} lg={12} xl={12}> <Col xs={24} sm={24} md={12} lg={12} xl={12}>
<Form.Checkbox <Form.Checkbox


+ 1
- 1
web/src/components/settings/personal/cards/NotificationSettings.jsx Переглянути файл

@@ -533,7 +533,7 @@ const NotificationSettings = ({
<CodeViewer <CodeViewer
content={{ content={{
type: 'quota_exceed', type: 'quota_exceed',
title: '额度预警通知',
title: t('额度预警通知'),
content: content:
'您的额度即将用尽,当前剩余额度为 {{value}}', '您的额度即将用尽,当前剩余额度为 {{value}}',
values: ['$0.99'], values: ['$0.99'],


Деякі файли не було показано, через те що забагато файлів було змінено

Завантаження…
Відмінити
Зберегти