62 Commits

Author SHA1 Message Date
  fengsilin caf457ce9e feat: 错误时记录上游响应体,支持流式 3 months ago
  fengsilin 575f84064e feat: codex API key 模式 + 前端凭证回显修复 3 months ago
  fengsilin ca91922904 docs: add log chat/upstream id design 3 months ago
  fengsilin 332d5a62f2 feat: support codex api key credentials 3 months ago
  fengsilin 1c0898257c feat: add redemption remarks 3 months ago
  fengsilin f18e51e3a9 feat: add redemption remark support 3 months ago
  fengsilin fae24fdf37 test: stabilize channel affinity usage cache tests 3 months ago
  fengsilin fc0257f6ca chore: ignore local worktrees 3 months ago
  fengsilin d8272e7707 feat(channel): 添加渠道"对外名称"(public_name)字段 3 months ago
  fengsilin 73d10b5799 feat: 错误日志记录上游 request-id 和响应体,Playground 渠道路由改为 header 传递 3 months ago
  fengsilin 5b8303ae81 fix: 登录后跳转来源页,充值金额校验优化 3 months ago
  fengsilin 8bb2883d7e feat(home): 首页定价卡片添加"去体验"按钮,跳转 Playground 3 months ago
  fengsilin a0ac37b2a9 fix: 替换 println 为 SysLog,修复 logger nil context 崩溃 3 months ago
  fengsilin 74cc8c0d56 feat(playground): 添加渠道选择功能,支持指定渠道体验模型 3 months ago
  fengsilin a34d44a824 fix(user): 修复从节点同步用户余额显示为 0 的问题 3 months ago
  fengsilin 967309fb56 feat: 添加邮箱后缀注册额度规则功能 3 months ago
  fengsilin 83f767b40b Merge branch 'worktree-feat-captcha' 3 months ago
  fengsilin d24d68a916 fix(i18n): 修复令牌和充值页面货币显示不一致问题 3 months ago
  fengsilin fe88bc7758 Merge branch 'worktree-feat-captcha' 3 months ago
  fengsilin 1912dd3b72 feat: 添加图片验证码功能,防止脚本批量注册 3 months ago
  fengsilin 20163dbad5 fix(legal): 法律文档页面切换语言时强制重新加载内容 3 months ago
  fengsilin 56e8d14485 fix: 语言切换后 TextArea 内容丢失 3 months ago
  fengsilin 65bd1dabb8 fix: 法律文档编辑器切换语言时 TextArea 内容未刷新 3 months ago
  fengsilin 934b4b10a7 feat: 法律文档双语配置(中/英)+ 精简语言支持至中英双语 3 months ago
  fengsilin 676902ed58 Merge branch 'worktree-feat-terms-usage-policy' 3 months ago
  fengsilin 2757f77011 feat(i18n): 补全服务条款和使用政策的多语言翻译 3 months ago
  fengsilin 65c6446a30 Merge branch 'worktree-feat-terms-usage-policy' 3 months ago
  fengsilin 1cc75aa9a2 feat: 服务条款和使用政策页面,支持后台 Markdown 配置 3 months ago
  fengsilin ed195cada5 fix(frontend): sort_order 默认值显示为"未设置",编辑弹窗增加提示 3 months ago
  fengsilin d5157a779c fix(pricing): 统一 sort_order 排序逻辑,有 meta 但未设置的不优先 3 months ago
  fengsilin 54e53108ec fix(sort): sort_order 默认值改为 999999,简化排序逻辑 3 months ago
  fengsilin 3b99c83e32 fix(sort): sort_order=0 的记录排到最后,非0值按升序排列 3 months ago
  fengsilin 518f1ab87f refactor(frontend): 移除拖拽排序,改为 sort_order 数值输入 3 months ago
  fengsilin 65880fc69b feat(frontend): Vendor Tab 和 Model 表格支持拖拽排序,移除定价页字母排序 3 months ago
  fengsilin 5e0f38d26c chore(frontend): 安装 @dnd-kit 拖拽排序库 3 months ago
  fengsilin e507897f21 feat(pricing): updatePricing 按 sort_order 排序 vendors 和 models 3 months ago
  fengsilin 749064932c feat: 新增 PUT /api/vendors/reorder 和 /api/models/reorder 批量排序接口 3 months ago
  fengsilin f8859b7c35 feat: Vendor 和 Model 添加 sort_order 字段,查询按 sort_order ASC 排序 3 months ago
  fengsilin 3a29f3772f feat(channel): 添加模型默认通道功能,支持优先路由和卡片标识 3 months ago
  fengsilin 93e331b624 feat: 支持 LOGO_FILE_PATH 环境变量指定本地 Logo,优化定价与语言设置 3 months ago
  fengsilin 94f24b29f7 style: 首页文案"全球"改为"顶级" 3 months ago
  fengsilin f0f0215b20 fix(pricing): 修复缓存倍率写入默认值1和缓存创建token丢失 3 months ago
  fengsilin 387f6c1ae0 feat(pricing): 定价数据源切换到渠道表,新增缓存价格展示 3 months ago
  fengsilin 00058cd671 refactor(channel-pricing): 提取辅助方法消除重复代码,修复缓存一致性 3 months ago
  fengsilin 42082311d3 fix: 修复默认语言不生效的问题 3 months ago
  fengsilin 9eb58c684e feat: 后台设置默认语言 3 months ago
  fengsilin a2b27ea2a6 docs: 后台设置默认语言实施计划 3 months ago
  fengsilin 436fcb405b docs: 后台设置默认语言功能设计文档 3 months ago
  fengsilin 8d22c4e80e style(home): 临时隐藏首页工具链/核心价值/工作流/生态伙伴 section 3 months ago
  fengsilin 355757404b merge: feat/channel-pricing-extended → master 3 months ago
  fengsilin 235f7c6e5f refactor(channel-pricing): 提取 ParseTagIds 辅助函数 + 移除调试 console.log 3 months ago
  fengsilin f3577590bc test(channel-pricing): 添加 API 端到端集成测试脚本 3 months ago
  fengsilin 129e17ac42 feat(channel-pricing): 前端支持缓存/图片/音频倍率编辑和展示 3 months ago
  fengsilin c84c26b5a2 feat(channel-pricing): GetChannelPricingByModelWithChannelInfo 支持 CASE WHEN 回退 + 扩展字段 3 months ago
  fengsilin d5d4714908 test(channel-pricing): 缓存写穿 + 字段默认值单元测试 3 months ago
  fengsilin 646f37dd4b feat(channel-pricing): Controller 扩展 API 支持新字段 + 输入校验 + 操作日志 3 months ago
  fengsilin 3a24fd598f feat(channel-pricing): ModelPriceHelper + UpdatePriceDataForChannelPricing 适配新签名,支持扩展比率覆盖 3 months ago
  fengsilin e5f029f91e feat(channel-pricing): InitDB 启动时全量加载渠道定价缓存 3 months ago
  fengsilin 2d4c73d3aa refactor(channel-pricing): 结构体新增扩展字段 + 缓存重写为全量加载写穿 3 months ago
  fengsilin 156618fdad merge: feat/alipay-payment → master 3 months ago
  fengsilin b514a0798b refactor(payment): 合并微信/支付宝支付重复代码,删除冗余文件 3 months ago
  fengsilin e8ade3abae feat(payment): 集成支付宝当面付扫码支付 + 补全前端 i18n 硬编码中文 3 months ago
100 changed files with 4624 additions and 510 deletions
Split View
  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 View File

@@ -7,4 +7,24 @@ Makefile
docs
.eslintcache
.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 View File

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

+ 1
- 0
.gitignore View File

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



+ 30
- 0
common/captcha.go View File

@@ -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 View File

@@ -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 View File

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

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

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

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


+ 18
- 0
common/str.go View File

@@ -87,6 +87,24 @@ func StringsContains(strs []string, str string) bool {
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
func StringToByteSlice(s string) []byte {
tmp1 := (*[2]uintptr)(unsafe.Pointer(&s))


+ 28
- 16
controller/channel.go View File

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

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

// 检查模型名称长度是否超过 255
for _, m := range channel.GetModels() {
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)
if channel.Type == constant.ChannelTypeCodex {
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")
}
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})
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())
c.JSON(http.StatusOK, gin.H{"success": false, "message": "刷新凭证失败,请稍后重试"})
return
@@ -2108,10 +2119,11 @@ func GetUserChannelsForBinding(c *gin.Context) {
result := make([]gin.H, 0, len(channels))
for _, ch := range channels {
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 View File

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

import (
"fmt"
"strconv"
"strings"

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

// CreateChannelPricingRequest 创建渠道定价请求
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 创建或更新渠道定价
@@ -68,38 +80,38 @@ func CreateChannelPricing(c *gin.Context) {
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)
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 {
common.ApiError(c, err)
return
}
common.SysLog(fmt.Sprintf("[ChannelPricing] updated: id=%d model=%s channel=%d", existing.Id, existing.ModelName, existing.ChannelId))
common.ApiSuccess(c, existing)
return
}

// 创建
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 {
common.ApiError(c, err)
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)
}

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

pricings := make([]*model.ChannelPricing, 0, len(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 在不同数据库表现不一致)
for _, cp := range pricings {
existing, _ := model.GetChannelPricing(cp.ModelName, cp.ChannelId)
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 {
common.ApiError(c, err)
return
@@ -168,6 +174,7 @@ func DeleteChannelPricing(c *gin.Context) {
return
}

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

@@ -204,6 +211,28 @@ func CopyGlobalPricing(c *gin.Context) {
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 ratio, completionRatio, price float64
@@ -224,28 +253,25 @@ func CopyGlobalPricing(c *gin.Context) {
}

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 {
imported++
}
} else {
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 {
imported++
}
}
}

common.SysLog(fmt.Sprintf("[ChannelPricing] copyGlobalPricing: channel=%d imported=%d/%d", channelId, imported, len(abilities)))
common.ApiSuccess(c, gin.H{
"total": len(abilities),
"imported": imported,
@@ -302,16 +328,7 @@ func GetChannelPricingWithTags(c *gin.Context) {
for _, cp := range list {
item := &ChannelPricingWithTags{
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)
}
@@ -323,3 +340,51 @@ func GetChannelPricingWithTags(c *gin.Context) {
"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 View File

@@ -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 View File

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

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

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

oauthKey, err := codex.ParseOAuthKey(strings.TrimSpace(ch.Key))
oauthKey, err := common.ParseCodexOAuthCredential(strings.TrimSpace(ch.Key))
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())
c.JSON(http.StatusOK, gin.H{"success": false, "message": "解析凭证失败,请检查渠道配置"})
return
}

accessToken := strings.TrimSpace(oauthKey.AccessToken)
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)
if err != nil {
@@ -98,6 +95,7 @@ func GetCodexChannelUsage(c *gin.Context) {

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

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


+ 128
- 0
controller/email_quota_rule.go View File

@@ -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 View File

@@ -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 View File

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

"usd_exchange_rate": operation_setting.USDExchangeRate,
"price": operation_setting.Price,
@@ -113,8 +115,10 @@ func GetStatus(c *gin.Context) {
"passkey_user_verification": passkeySetting.UserVerification,
"passkey_attachment": passkeySetting.AttachmentPreference,
"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,
"_qn": "new-api",
}
@@ -188,20 +192,53 @@ func GetAbout(c *gin.Context) {
return
}

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

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

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{
"success": true,
"message": "",
"data": system_setting.GetLegalSettings().PrivacyPolicy,
"data": getLegalContent(ls.UsagePolicyZh, ls.UsagePolicyEn, lang),
})
return
}
@@ -228,7 +265,39 @@ func GetHomePageContent(c *gin.Context) {
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) {
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")
if err := common.Validate.Var(email, "required,email"); err != nil {
c.JSON(http.StatusOK, gin.H{


+ 62
- 0
controller/playground_channels.go View File

@@ -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 View File

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


+ 159
- 0
controller/redemption_test.go View File

@@ -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 View File

@@ -369,6 +369,12 @@ func processChannelError(c *gin.Context, channelError types.ChannelError, err *t
other["channel_id"] = channelId
other["channel_name"] = c.GetString("channel_name")
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["use_channel"] = c.GetStringSlice("use_channel")
isMultiKey := common.GetContextKeyBool(c, constant.ContextKeyChannelIsMultiKey)


+ 55
- 0
controller/reorder.go View File

@@ -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 View File

@@ -72,17 +72,38 @@ func GetTopUpInfo(c *gin.Context) {
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{
"enable_online_topup": enableOnlineTopup,
"enable_stripe_topup": setting.StripeApiSecret != "" && setting.StripeWebhookSecret != "" && setting.StripePriceId != "",
"enable_creem_topup": setting.CreemApiKey != "" && setting.CreemProducts != "[]",
"enable_wechat_topup": setting.IsWechatPayConfigured(),
"enable_alipay_topup": setting.IsAlipayConfigured(),
"creem_products": setting.CreemProducts,
"pay_methods": payMethods,
"min_topup": operation_setting.MinTopUp,
"stripe_min_topup": setting.StripeMinTopUp,
"wechat_pay_min_topup": setting.WechatPayMinTopUp,
"alipay_pay_min_topup": setting.AlipayMinTopUp,
"amount_options": operation_setting.GetPaymentSetting().AmountOptions,
"discount": operation_setting.GetPaymentSetting().AmountDiscount,
}
@@ -143,15 +164,37 @@ func getPayMoney(amount int64, group string) float64 {
}

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 {
dMinTopup := decimal.NewFromInt(int64(minTopup))
dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit)
minTopup = int(dMinTopup.Mul(dQuotaPerUnit).IntPart())
minTopup = minTopup * int(common.QuotaPerUnit)
}
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) {
var req EpayRequest
err := c.ShouldBindJSON(&req)


+ 280
- 0
controller/topup_alipay.go View File

@@ -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 View File

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

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

// getWechatPayMoney 计算微信支付应付金额(元)
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 {
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 View File

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


+ 1
- 0
docs/DATABASE_SCHEMA.md View File

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


+ 307
- 0
docs/superpowers/plans/2026-04-17-default-language-setting.md View File

@@ -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 View File

@@ -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 View File

@@ -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 View File

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


+ 2
- 0
go.mod View File

@@ -94,6 +94,7 @@ require (
github.com/go-sql-driver/mysql v1.7.0 // indirect
github.com/go-webauthn/x v0.1.25 // 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/gorilla/context 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/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // 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/ncruces/go-strftime v0.1.9 // indirect
github.com/pelletier/go-toml/v2 v2.2.1 // indirect


+ 57
- 0
go.sum View File

@@ -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/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/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/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.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.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/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
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 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
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/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
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/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/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/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
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=
golang.org/x/arch v0.21.0 h1:iTC9o7+wP6cPWpDWkivCvQFGAHDQ59SrSxsLPcnkArw=
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-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/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/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8=
golang.org/x/image v0.23.0 h1:HseQ7c2OpPKTPVzNjG5fwJsOTCiiwS4QdsYi5XU6H68=
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/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-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/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/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-20190916202348-b4ddaad3f8a3/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-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-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.8.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/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-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.3/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/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-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/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=
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=


+ 5
- 3
logger/logger.go View File

@@ -79,9 +79,11 @@ func logHelper(ctx context.Context, level string, msg string) {
if level == loggerINFO {
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()
_, _ = fmt.Fprintf(writer, "[%s] %v | %s | %s \n", level, now.Format("2006/01/02 - 15:04:05"), id, msg)


+ 10
- 0
main.go View File

@@ -137,6 +137,16 @@ func main() {
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" {
gopool.Go(func() {
log.Println(http.ListenAndServe("0.0.0.0:8005", nil))


+ 62
- 14
middleware/distributor.go View File

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

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) {
return func(c *gin.Context) {
var channel *model.Channel
channelId, ok := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId)
modelRequest, shouldSelectChannel, err := getModelRequest(c)
if err != nil {
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()}))
return
}
channelId, ok := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId)
if ok {
id, err := strconv.Atoi(channelId.(string))
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)
autoGroups := service.GetUserAutoGroup(userGroup)
for _, g := range autoGroups {
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) {
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, defaultCh.Id) {
channel = defaultCh
selectGroup = 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
}
}
} 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
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
}
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)
// 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 != "" {


+ 53
- 0
model/ability.go View File

@@ -66,6 +66,59 @@ func GetAbilitiesByChannelId(channelId int) ([]*Ability, error) {
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) {

var priorities []int


+ 9
- 1
model/channel.go View File

@@ -26,6 +26,7 @@ type Channel struct {
TestModel *string `json:"test_model"`
Status int `json:"status" gorm:"default:1"`
Name string `json:"name" gorm:"index"`
PublicName string `json:"public_name" gorm:"size:255;default:''"`
Weight *uint `json:"weight" gorm:"default:0"`
CreatedTime int64 `json:"created_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)
}

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

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


+ 226
- 79
model/channel_pricing.go View File

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

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

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

// QuotaType 计费类型
@@ -41,6 +42,40 @@ type ChannelPricing struct {
CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
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 {
@@ -49,31 +84,43 @@ func (cp *ChannelPricing) Insert() error {
cp.UpdatedTime = now
err := DB.Create(cp).Error
if err == nil {
InvalidateChannelPricingCache()
setCache(getChannelPricingCacheKey(cp.ModelName, cp.ChannelId), cp)
if cp.IsDefault {
setDefaultChannelCache(cp.ModelName, cp.ChannelId)
}
}
return err
}

func (cp *ChannelPricing) Update() error {
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 {
InvalidateChannelPricingCache()
setCache(getChannelPricingCacheKey(cp.ModelName, cp.ChannelId), cp)
if cp.IsDefault {
setDefaultChannelCache(cp.ModelName, cp.ChannelId)
} else {
clearDefaultChannelCacheIfMatch(cp.ModelName, cp.Id)
}
}
return err
}

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

// getChannelPricingCacheKey 生成缓存键
@@ -142,71 +200,51 @@ func getChannelPricingCacheKey(modelName string, channelId int) string {
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()
// 检查缓存是否过期
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()

// 缓存未命中或已过期,查询数据库
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
if err := DB.Find(&pricings).Error; err != nil {
common.SysError("[ChannelPricing] LoadChannelPricingCache failed: " + err.Error())
return
}

channelPricingCacheLock.Lock()
channelPricingCache = make(map[string]*ChannelPricing, len(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 带渠道信息的定价响应
@@ -214,6 +252,7 @@ type ChannelPricingWithChannel struct {
Id int `json:"id"`
ChannelId int `json:"channel_id"`
ChannelName string `json:"channel_name"`
ChannelPublicName string `json:"channel_public_name"`
ChannelType int `json:"channel_type"`
TagIds string `json:"tag_ids" gorm:"column:tag_ids"` // 渠道定价的标签ID列表(逗号分隔)
Tags []*PricingTag `json:"tags" gorm:"-"` // 渠道定价的标签详情(不参与数据库扫描)
@@ -221,7 +260,13 @@ type ChannelPricingWithChannel struct {
ModelRatio float64 `json:"model_ratio"`
CompletionRatio float64 `json:"completion_ratio"`
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 获取指定模型的渠道定价(带渠道信息)
@@ -249,16 +294,22 @@ func GetChannelPricingByModelWithChannelInfo(modelName string) ([]*ChannelPricin
if !hasPrice {
globalModelPrice = 0
}

// 查询所有支持该模型的渠道,左连接渠道定价表
// 高级字段(cache/image/audio)不回退全局值,直接返回 0
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.model_ratio, ?) as model_ratio,
COALESCE(channel_pricings.completion_ratio, ?) as completion_ratio,
COALESCE(channel_pricings.model_price, ?) as model_price,
channel_pricings.id as id,
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`,
defaultQuotaType, globalModelRatio, globalCompletionRatio, globalModelPrice).
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.enabled = ?", true).
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
if err != nil {
return nil, err
@@ -286,17 +337,113 @@ func GetChannelPricingByModelWithChannelInfo(modelName string) ([]*ChannelPricin

// 为每个渠道定价填充标签
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 View File

@@ -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 View File

@@ -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 View File

@@ -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 View File

@@ -203,7 +203,12 @@ func InitDB() (err error) {
}
common.SysLog("database migration started")
err = migrateDB()
return err
if err != nil {
return err
}
LoadEmailQuotaCache()
LoadChannelPricingCache()
return nil
} else {
common.FatalLog(err)
}
@@ -282,6 +287,7 @@ func migrateDB() error {
&PricingTag{},
&PendingSyncRecord{},
&QuotaSyncLog{},
&EmailQuotaRule{},
)
if err != nil {
return err
@@ -294,9 +300,14 @@ func migrateDB() error {
if err := DB.AutoMigrate(&SubscriptionPlan{}); err != nil {
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 {
// Drop bound_channel_id column from tokens table (deprecated field)
@@ -336,6 +347,7 @@ func migrateDBFast() error {
{&PricingTag{}, "PricingTag"},
{&PendingSyncRecord{}, "PendingSyncRecord"},
{&QuotaSyncLog{}, "QuotaSyncLog"},
{&EmailQuotaRule{}, "EmailQuotaRule"},
}
// 动态计算migration数量,确保errChan缓冲区足够大
errChan := make(chan error, len(migrations))
@@ -683,3 +695,14 @@ func PingDB() error {
common.SysLog("Database pinged successfully")
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 View File

@@ -42,6 +42,7 @@ type Model struct {
Endpoints string `json:"endpoints,omitempty" gorm:"type:text"`
Status int `json:"status" 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"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
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()
// 使用 Select 强制更新所有字段,包括零值
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
}

@@ -97,6 +98,16 @@ func (mi *Model) Delete() 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) {
var stats []struct {
VendorID int64
@@ -117,7 +128,7 @@ func GetVendorModelCounts() (map[int64]int64, error) {

func GetAllModels(offset int, limit int) ([]*Model, error) {
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
}

@@ -165,8 +176,23 @@ func SearchModels(keyword string, vendor string, offset int, limit int) ([]*Mode
if err := db.Count(&total).Error; err != nil {
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 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 View File

@@ -43,6 +43,7 @@ func InitOptionMap() {
common.OptionMap["TelegramOAuthEnabled"] = strconv.FormatBool(common.TelegramOAuthEnabled)
common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled)
common.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(common.TurnstileCheckEnabled)
common.OptionMap["CaptchaEnabled"] = strconv.FormatBool(common.CaptchaEnabled)
common.OptionMap["RegisterEnabled"] = strconv.FormatBool(common.RegisterEnabled)
common.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(common.AutomaticDisableChannelEnabled)
common.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(common.AutomaticEnableChannelEnabled)
@@ -101,6 +102,12 @@ func InitOptionMap() {
common.OptionMap["WechatPayPubKeyB64"] = setting.WechatPayPubKeyB64
common.OptionMap["WechatPayMinTopUp"] = strconv.Itoa(setting.WechatPayMinTopUp)
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["Chats"] = setting.Chats2JsonString()
common.OptionMap["AutoGroups"] = setting.AutoGroups2JsonString()
@@ -178,6 +185,13 @@ func triggerWechatPayReset() {
}
}

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

func loadOptionsFromDatabase() {
options, _ := AllOption()
for _, option := range options {
@@ -255,6 +269,8 @@ func updateOptionMap(key string, value string) (err error) {
common.TelegramOAuthEnabled = boolValue
case "TurnstileCheckEnabled":
common.TurnstileCheckEnabled = boolValue
case "CaptchaEnabled":
common.CaptchaEnabled = boolValue
case "RegisterEnabled":
common.RegisterEnabled = boolValue
case "EmailDomainRestrictionEnabled":
@@ -410,6 +426,21 @@ func updateOptionMap(key string, value string) (err error) {
setting.WechatPayMinTopUp, _ = strconv.Atoi(value)
case "WechatPayUnitPrice":
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":
err = common.UpdateTopupGroupRatioByJSONString(value)
case "GitHubClientId":
@@ -428,6 +459,8 @@ func updateOptionMap(key string, value string) (err error) {
common.SystemName = value
case "Logo":
common.Logo = value
case "DefaultLanguage":
common.DefaultLanguage = value
case "WeChatServerAddress":
common.WeChatServerAddress = value
case "WeChatServerToken":


+ 134
- 9
model/pricing.go View File

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

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

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

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

vendorsList = make([]PricingVendor, 0, len(vendorMap))
for _, v := range vendorMap {
vendorsList = append(vendorsList, PricingVendor{
@@ -171,6 +180,14 @@ func updatePricing() {
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])

@@ -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)
for model, groups := range modelGroupsMap {
pricing := Pricing{
@@ -289,19 +324,37 @@ func updatePricing() {
pricing.VendorID = meta.VendorID
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)
}

// 按 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 {
pricingMap[0].PricingVersion = "82c4a357505fff6fee8462c3f7ec8a645bb95532669cb73b2cabee6a416ec24f"
@@ -324,3 +377,75 @@ func updatePricing() {
func GetSupportedEndpointMap() map[string]common.EndpointInfo {
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 View File

@@ -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 View File

@@ -20,6 +20,7 @@ type Redemption struct {
Key string `json:"key" gorm:"type:char(32);uniqueIndex"`
Status int `json:"status" gorm:"default:1"`
Name string `json:"name" gorm:"index"`
Remark string `json:"remark" gorm:"index"`
Quota int `json:"quota" gorm:"default:100"`
CreatedTime int64 `json:"created_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
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 {
query = query.Where("name LIKE ?", keyword+"%")
query = query.Where("name LIKE ? OR remark LIKE ?", keyword+"%", keyword+"%")
}

// 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
func (redemption *Redemption) Update() 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
}



+ 101
- 14
model/redemption_test.go View File

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

// 创建兑换码
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)

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

// 创建已过期的兑换码
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)

@@ -252,3 +252,90 @@ func TestRedeem_SyncedUser(t *testing.T) {
require.NoError(t, db.First(&updatedUser, 100).Error)
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 View File

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

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

"github.com/shopspring/decimal"
@@ -21,6 +22,32 @@ type TopUp struct {
CreateTime int64 `json:"create_time"`
CompleteTime int64 `json:"complete_time"`
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 {
@@ -135,6 +162,7 @@ func GetUserTopUps(userId int, pageInfo *common.PageInfo) (topups []*TopUp, tota
return nil, 0, err
}

fillTopUpEmails(topups)
return topups, total, nil
}

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

fillTopUpEmails(topups)
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 {
return nil, 0, err
}
fillTopUpEmails(topups)
return topups, total, nil
}

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

@@ -269,7 +300,7 @@ func ManualCompleteTopUp(tradeNo string) error {
// 计算应充值额度:
// - Stripe/微信支付订单:Money 代表经分组倍率换算后的数量,直接 * 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)
quotaToAdd = int(decimal.NewFromFloat(topUp.Money).Mul(dQuotaPerUnit).IntPart())
} else {


+ 12
- 6
model/topup_wechat.go View File

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

// RechargeWechat 微信支付充值完成(由回调触发)
// 与 Recharge/RechargeCreem 类似,使用事务+行锁保证幂等
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 == "" {
return errors.New("未提供支付单号")
}
@@ -34,7 +44,6 @@ func RechargeWechat(tradeNo string) error {
}

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

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

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

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


+ 40
- 6
model/user.go View File

@@ -214,6 +214,20 @@ func GetMaxUserId() int {
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) {
// Start transaction
tx := DB.Begin()
@@ -245,6 +259,7 @@ func GetAllUsers(pageInfo *common.PageInfo) (users []*User, total int64, err err
return nil, 0, err
}

applySyncedUserQuota(users)
return users, total, nil
}

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

applySyncedUserQuota(users)
return users, total, nil
}

@@ -410,7 +426,12 @@ func (user *User) Insert(inviterId int) error {
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.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 common.QuotaForInvitee > 0 {
@@ -469,7 +494,12 @@ func (user *User) InsertWithTx(tx *gorm.DB, inviterId int) error {
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)

// 初始化用户设置
@@ -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 common.QuotaForInvitee > 0 {


+ 18
- 2
model/vendor_meta.go View File

@@ -18,6 +18,7 @@ type Vendor struct {
Description string `json:"description,omitempty" gorm:"type:text"`
Icon string `json:"icon,omitempty" gorm:"type:varchar(128)"`
Status int `json:"status" gorm:"default:1"`
SortOrder int `json:"sort_order" gorm:"default:999999"`
CreatedTime int64 `json:"created_time" gorm:"bigint"`
UpdatedTime int64 `json:"updated_time" gorm:"bigint"`
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 获取全部供应商(分页)
func GetAllVendors(offset int, limit int) ([]*Vendor, error) {
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
}

@@ -81,8 +82,23 @@ func SearchVendors(keyword string, offset int, limit int) ([]*Vendor, int64, err
return nil, 0, err
}
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 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 View File

@@ -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)
}
if common2.DebugEnabled {
println("fullRequestURL:", fullRequestURL)
common2.SysLog("fullRequestURL: " + fullRequestURL)
}
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
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)
}
if common2.DebugEnabled {
println("fullRequestURL:", fullRequestURL)
common2.SysLog("fullRequestURL: " + fullRequestURL)
}
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
if err != nil {


+ 4
- 4
relay/channel/claude/relay-claude.go View File

@@ -635,7 +635,7 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d
if claudeResponse.Message != nil && claudeResponse.Message.Usage != nil {
claudeInfo.Usage.PromptTokens = claudeResponse.Message.Usage.InputTokens
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.ClaudeCacheCreation1hTokens = claudeResponse.Message.Usage.GetCacheCreation1hTokens()
claudeInfo.Usage.CompletionTokens = claudeResponse.Message.Usage.OutputTokens
@@ -659,8 +659,8 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d
if claudeResponse.Usage.CacheReadInputTokens > 0 {
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 {
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.TotalTokens = claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens
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.ClaudeCacheCreation1hTokens = claudeResponse.Usage.GetCacheCreation1hTokens()
}


+ 27
- 11
relay/channel/codex/adaptor.go View File

@@ -138,9 +138,21 @@ func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
if info.RelayMode != relayconstant.RelayModeResponses && info.RelayMode != relayconstant.RelayModeResponsesCompact {
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 {
path = "/backend-api/codex/responses/compact"
path = "/v1/responses/compact"
}
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)

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 {
return err
}
@@ -178,13 +199,8 @@ func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Header, info *rel
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")
if info.IsStream {
req.Set("Accept", "text/event-stream")
} else if req.Get("Accept") == "" {
if req.Get("Accept") == "" {
req.Set("Accept", "application/json")
}



+ 5
- 9
relay/channel/openai/chat_via_responses.go View File

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

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)
@@ -484,14 +486,8 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
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

default:


+ 11
- 4
relay/channel/openai/relay_responses.go View File

@@ -31,7 +31,9 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
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() {
@@ -78,10 +80,10 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp

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

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

// 检查当前数据是否包含 completed 状态和 usage 信息
var streamResponse dto.ResponsesStreamResponse
if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil {
sendResponsesStreamData(c, streamResponse, data)
@@ -109,10 +111,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
}
}
case "response.output_text.delta":
// 处理输出文本
responseTextBuilder.WriteString(streamResponse.Delta)
case dto.ResponsesOutputTypeItemDone:
// 函数调用处理
if streamResponse.Item != nil {
switch streamResponse.Item.Type {
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 {
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
})

if streamErr != nil {
return nil, streamErr
}

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


+ 29
- 0
relay/channel/openai/responses_error.go View File

@@ -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 View File

@@ -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 View File

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

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


+ 17
- 4
relay/compatible_handler.go View File

@@ -27,6 +27,22 @@ import (
"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) {
info.InitChannelMeta(c)

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

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)
usage, newApiErr := chatCompletionsViaResponses(c, info, adaptor, request)
if newApiErr != nil {


+ 60
- 41
relay/helper/price.go View File

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

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
@@ -14,9 +15,6 @@ import (
"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
func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) types.GroupRatioInfo {
groupRatioInfo := types.GroupRatioInfo{
@@ -52,17 +50,30 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
var modelRatio float64
var completionRatio float64
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
// 渠道定价的扩展比率(非零值直接使用)
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)

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
if !usePrice {
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)
}
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
preConsumedQuota = int(float64(preConsumedTokens) * ratio)
} else {
@@ -120,6 +116,24 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
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
if !operation_setting.GetQuotaSetting().EnableFreeModelPreConsume {
// 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,
GroupRatioInfo: groupRatioInfo,
UsePrice: usePrice,
QuotaToPreConsume: preConsumedQuota,
CacheRatio: cacheRatio,
ImageRatio: imageRatio,
AudioRatio: audioRatio,
AudioCompletionRatio: audioCompletionRatio,
CacheCreationRatio: cacheCreationRatio,
CacheCreation5mRatio: cacheCreationRatio5m,
CacheCreation1hRatio: cacheCreationRatio1h,
QuotaToPreConsume: preConsumedQuota,
CacheCreation5mRatio: cacheCreationRatio,
CacheCreation1hRatio: cacheCreationRatio * types.ClaudeCacheCreation1hMultiplier,
}

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

cpRatio, cpCompletionRatio, cpPrice, cpUsePrice, found := model.GetEffectivePricing(info.OriginModelName, channelId)
cp, found := model.GetEffectivePricing(info.OriginModelName, channelId)
if !found {
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 {
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 {
estimateTokens := info.GetEstimatePromptTokens()
if estimateTokens > 0 {
ratio := cpRatio * info.PriceData.GroupRatioInfo.GroupRatio
ratio := cp.ModelRatio * info.PriceData.GroupRatioInfo.GroupRatio
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 View File

@@ -26,11 +26,14 @@ func SetApiRouter(router *gin.Engine) {
apiRouter.GET("/notice", controller.GetNotice)
apiRouter.GET("/user-agreement", controller.GetUserAgreement)
apiRouter.GET("/privacy-policy", controller.GetPrivacyPolicy)
apiRouter.GET("/terms", controller.GetTermsOfService)
apiRouter.GET("/usage-policy", controller.GetUsagePolicy)
apiRouter.GET("/about", controller.GetAbout)
//apiRouter.GET("/midjourney", controller.GetMidjourney)
apiRouter.GET("/home_page_content", controller.GetHomePageContent)
apiRouter.GET("/pricing", middleware.TryUserAuth(), controller.GetPricing)
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("/reset_password", middleware.CriticalRateLimit(), middleware.TurnstileCheck(), controller.SendPasswordResetEmail)
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("/creem/webhook", controller.CreemWebhook)
apiRouter.POST("/wechat/pay/webhook", controller.WechatPayWebhook)
apiRouter.POST("/alipay/pay/webhook", controller.AlipayPayWebhook)
// Universal secure verification routes
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", controller.GetSelf)
selfRoute.GET("/models", controller.GetUserModels)
selfRoute.GET("/model_channels", controller.GetModelChannels)
selfRoute.GET("/channels", controller.GetUserChannelsForBinding)
selfRoute.PUT("/self", controller.UpdateSelf)
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", controller.RequestWechatPay)
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.PUT("/setting", controller.UpdateUserSetting)

@@ -190,6 +198,8 @@ func SetApiRouter(router *gin.Engine) {
channelPricingRoute.POST("/batch", controller.BatchCreateChannelPricing)
channelPricingRoute.POST("/copy_global/:channel_id", controller.CopyGlobalPricing)
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)
}

// 邮箱后缀额度规则路由(管理员权限)
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)
customOAuthRoute := apiRouter.Group("/custom-oauth-provider")
customOAuthRoute.Use(middleware.RootAuth())
@@ -350,6 +370,7 @@ func SetApiRouter(router *gin.Engine) {
vendorRoute.GET("/:id", controller.GetVendorMeta)
vendorRoute.POST("/", controller.CreateVendorMeta)
vendorRoute.PUT("/", controller.UpdateVendorMeta)
vendorRoute.PUT("/reorder", controller.ReorderVendors)
vendorRoute.DELETE("/:id", controller.DeleteVendorMeta)
}

@@ -364,6 +385,7 @@ func SetApiRouter(router *gin.Engine) {
modelsRoute.GET("/:id", controller.GetModelMeta)
modelsRoute.POST("/", controller.CreateModelMeta)
modelsRoute.PUT("/", controller.UpdateModelMeta)
modelsRoute.PUT("/reorder", controller.ReorderModels)
modelsRoute.DELETE("/:id", controller.DeleteModelMeta)
}



+ 8
- 0
router/web-router.go View File

@@ -14,6 +14,14 @@ import (
)

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(middleware.GlobalWebRateLimit())
router.Use(middleware.Cache())


+ 6
- 7
service/channel_affinity_usage_cache_test.go View File

@@ -4,7 +4,6 @@ import (
"fmt"
"net/http/httptest"
"testing"
"time"

"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/types"
@@ -26,9 +25,9 @@ func buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP string)
}

func TestObserveChannelAffinityUsageCacheByRelayFormat_ClaudeMode(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)

usage := &dto.Usage{
@@ -53,9 +52,9 @@ func TestObserveChannelAffinityUsageCacheByRelayFormat_ClaudeMode(t *testing.T)
}

func TestObserveChannelAffinityUsageCacheByRelayFormat_MixedMode(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)

openAIUsage := &dto.Usage{
@@ -83,9 +82,9 @@ func TestObserveChannelAffinityUsageCacheByRelayFormat_MixedMode(t *testing.T) {
}

func TestObserveChannelAffinityUsageCacheByRelayFormat_UnsupportedModeKeepsEmpty(t *testing.T) {
ruleName := fmt.Sprintf("rule_%d", time.Now().UnixNano())
ruleName := "rule_" + t.Name()
usingGroup := "default"
keyFP := fmt.Sprintf("fp_%d", time.Now().UnixNano())
keyFP := "fp_" + t.Name()
ctx := buildChannelAffinityStatsContextForTest(ruleName, usingGroup, keyFP)

usage := &dto.Usage{


+ 2
- 24
service/codex_credential_refresh.go View File

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

import (
"context"
"errors"
"fmt"
"strings"
"time"
@@ -16,28 +15,7 @@ type CodexCredentialRefreshOptions struct {
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) {
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")
}

oauthKey, err := parseCodexOAuthKey(strings.TrimSpace(ch.Key))
oauthKey, err := common.ParseCodexOAuthCredential(strings.TrimSpace(ch.Key))
if err != nil {
return nil, nil, err
}


+ 1
- 1
service/codex_credential_refresh_task.go View File

@@ -93,7 +93,7 @@ func runCodexCredentialAutoRefreshOnce() {
continue
}

oauthKey, err := parseCodexOAuthKey(rawKey)
oauthKey, err := common.ParseCodexOAuthCredential(rawKey)
if err != nil {
continue
}


+ 21
- 0
service/error.go View File

@@ -58,6 +58,14 @@ func MidjourneyErrorWithStatusCodeWrapper(code int, desc string, statusCode int)
// 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 {
text := err.Error()
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) {
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)
if err != nil {
return
}
CloseResponseBodyGracefully(resp)
bodyStr := TruncateBody(string(responseBody))
newApiErr.UpstreamBody = bodyStr
var errResponse dto.GeneralErrorResponse
buildErrWithBody := func(message string) error {
if message == "" {
@@ -115,6 +132,8 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai
oaiError := errResponse.TryToOpenAIError()
if oaiError != nil {
newApiErr = types.WithOpenAIError(*oaiError, resp.StatusCode)
newApiErr.UpstreamRequestId = upstreamReqId
newApiErr.UpstreamBody = bodyStr
if showBodyWhenFail {
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.UpstreamRequestId = upstreamReqId
newApiErr.UpstreamBody = bodyStr
if showBodyWhenFail {
newApiErr.Err = buildErrWithBody(newApiErr.Error())
}


+ 32
- 0
service/truncate_body_test.go View File

@@ -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 View File

@@ -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 View File

@@ -604,6 +604,15 @@ func GetAudioRatio(name string) float64 {
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 {
name = FormatMatchingModelName(name)
if ratio, ok := audioCompletionRatioMap.Get(name); ok {
@@ -612,6 +621,15 @@ func GetAudioCompletionRatio(name string) float64 {
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 {
name = FormatMatchingModelName(name)
_, ok := audioRatioMap.Get(name)


+ 16
- 4
setting/system_setting/legal.go View File

@@ -3,13 +3,25 @@ package system_setting
import "github.com/QuantumNous/new-api/setting/config"

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{
UserAgreement: "",
PrivacyPolicy: "",
UserAgreementZh: "",
UserAgreementEn: "",
PrivacyPolicyZh: "",
PrivacyPolicyEn: "",
TermsOfServiceZh: "",
TermsOfServiceEn: "",
UsagePolicyZh: "",
UsagePolicyEn: "",
}

func init() {


+ 122
- 0
test-scripts/test_channel_pricing_extended.sh View File

@@ -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 View File

@@ -88,14 +88,16 @@ const (
)

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.


+ 26
- 0
types/price_data.go View File

@@ -2,6 +2,10 @@ package types

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 {
GroupRatio float64
GroupSpecialRatio float64
@@ -27,6 +31,28 @@ type PriceData struct {
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) {
if p.OtherRatios == nil {
p.OtherRatios = make(map[string]float64)


+ 1
- 1
web/index.html View File

@@ -10,7 +10,7 @@
content="OpenAI 接口聚合管理,支持多种渠道包括 Azure,可用于二次分发管理 key,仅单可执行文件,已打包好 Docker 镜像,一键部署,开箱即用"
/>
<meta name="generator" content="new-api" />
<title>New API</title>
<title>Loading...</title>
<!--umami-->
<!--Google Analytics-->
</head>


+ 18
- 0
web/src/App.jsx View File

@@ -55,6 +55,8 @@ const Dashboard = lazy(() => import('./pages/Dashboard'));
const About = lazy(() => import('./pages/About'));
const UserAgreement = lazy(() => import('./pages/UserAgreement'));
const PrivacyPolicy = lazy(() => import('./pages/PrivacyPolicy'));
const Terms = lazy(() => import('./pages/Terms'));
const UsagePolicy = lazy(() => import('./pages/UsagePolicy'));

function DynamicOAuth2Callback() {
const { provider } = useParams();
@@ -358,6 +360,22 @@ function App() {
</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
path='/console/chat/:id?'
element={


+ 21
- 11
web/src/components/auth/LoginForm.jsx View File

@@ -17,8 +17,8 @@ along with this program. If not, see <https://www.gnu.org/licenses/>.
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 { StatusContext } from '../../context/Status';
import {
@@ -69,6 +69,7 @@ import { SiDiscord } from 'react-icons/si';

const LoginForm = () => {
let navigate = useNavigate();
const location = useLocation();
const { t } = useTranslation();
const githubButtonTextKeyByState = {
idle: '使用 GitHub 继续',
@@ -113,6 +114,15 @@ const LoginForm = () => {
const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]);
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 systemName = getSystemName();

@@ -198,7 +208,7 @@ const LoginForm = () => {
localStorage.setItem('user', JSON.stringify(data));
setUserData(data);
updateAPI();
navigate('/');
navigateAfterLogin();
showSuccess(t('登录成功!'));
setShowWeChatLoginModal(false);
} else {
@@ -255,7 +265,7 @@ const LoginForm = () => {
centered: true,
});
}
navigate('/console');
navigateAfterLogin();
} else {
showError(message);
}
@@ -300,7 +310,7 @@ const LoginForm = () => {
showSuccess(t('登录成功!'));
setUserData(data);
updateAPI();
navigate('/');
navigateAfterLogin();
} else {
showError(message);
}
@@ -456,7 +466,7 @@ const LoginForm = () => {
setUserData(finish.data);
updateAPI();
showSuccess(t('登录成功!'));
navigate('/console');
navigateAfterLogin();
} else {
showError(finish.message || t('Passkey 登录失败,请重试'));
}
@@ -490,8 +500,8 @@ const LoginForm = () => {
userDispatch({ type: 'login', payload: data });
setUserData(data);
updateAPI();
showSuccess('登录成功!');
navigate('/console');
showSuccess(t('登录成功!'));
navigateAfterLogin();
};

// 返回登录页面
@@ -505,7 +515,7 @@ const LoginForm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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'>
{systemName}
</Title>
@@ -721,7 +731,7 @@ const LoginForm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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>
</div>

@@ -885,7 +895,7 @@ const LoginForm = () => {
}}
>
<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 className='text-center mb-4'>


+ 1
- 1
web/src/components/auth/PasswordResetConfirm.jsx View File

@@ -118,7 +118,7 @@ const PasswordResetConfirm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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'>
{systemName}
</Title>


+ 1
- 1
web/src/components/auth/PasswordResetForm.jsx View File

@@ -118,7 +118,7 @@ const PasswordResetForm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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'>
{systemName}
</Title>


+ 68
- 9
web/src/components/auth/RegisterForm.jsx View File

@@ -108,6 +108,9 @@ const RegisterForm = () => {
const [hasPrivacyPolicy, setHasPrivacyPolicy] = useState(false);
const [githubButtonState, setGithubButtonState] = useState('idle');
const [githubButtonDisabled, setGithubButtonDisabled] = useState(false);
const [captchaId, setCaptchaId] = useState('');
const [captchaImage, setCaptchaImage] = useState('');
const [captchaCode, setCaptchaCode] = useState('');
const githubTimeoutRef = useRef(null);
const githubButtonText = t(githubButtonTextKeyByState[githubButtonState]);

@@ -141,6 +144,8 @@ const RegisterForm = () => {
hasCustomOAuthProviders,
);

const captchaEnabled = !!status?.captcha_enabled;

const [showEmailVerification, setShowEmailVerification] = useState(false);

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 = () => {
setWechatLoading(true);
setShowWeChatLoginModal(true);
@@ -184,7 +209,7 @@ const RegisterForm = () => {

const onSubmitWeChatVerificationCode = async () => {
if (turnstileEnabled && turnstileToken === '') {
showInfo('请稍后几秒重试,Turnstile 正在检查用户环境!');
showInfo(t('请稍后几秒重试,Turnstile 正在检查用户环境!'));
return;
}
setWechatCodeSubmitLoading(true);
@@ -257,21 +282,28 @@ const RegisterForm = () => {
const sendVerificationCode = async () => {
if (inputs.email === '') return;
if (turnstileEnabled && turnstileToken === '') {
showInfo('请稍后几秒重试,Turnstile 正在检查用户环境!');
showInfo(t('请稍后几秒重试,Turnstile 正在检查用户环境!'));
return;
}
if (captchaEnabled && captchaCode === '') {
showInfo(t('请先输入图片验证码'));
return;
}
setVerificationCodeLoading(true);
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;
if (success) {
showSuccess(t('验证码发送成功,请检查你的邮箱!'));
setDisableButton(true); // 发送成功后禁用按钮,开始倒计时
setDisableButton(true);
} else {
showError(message);
}
if (captchaEnabled) loadCaptcha();
} catch (error) {
showError(t('发送验证码失败,请重试'));
} finally {
@@ -396,7 +428,7 @@ const RegisterForm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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'>
{systemName}
</Title>
@@ -559,7 +591,7 @@ const RegisterForm = () => {
<div className='flex flex-col items-center'>
<div className='w-full max-w-md'>
<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'>
{systemName}
</Title>
@@ -624,6 +656,33 @@ const RegisterForm = () => {
</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
field='verification_code'
label={t('验证码')}
@@ -745,7 +804,7 @@ const RegisterForm = () => {
}}
>
<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 className='text-center mb-4'>


+ 3
- 2
web/src/components/common/markdown/MarkdownRenderer.jsx View File

@@ -268,7 +268,7 @@ export function PreCode(props) {
color: 'var(--semi-color-text-2)',
}}
>
HTML预览:
{t('HTML预览:')}
</div>
<SandboxedHtmlPreview code={htmlCode} />
</div>
@@ -635,6 +635,7 @@ function _MarkdownContent(props) {
export const MarkdownContent = React.memo(_MarkdownContent);

export function MarkdownRenderer(props) {
const { t } = useTranslation();
const {
content,
loading,
@@ -680,7 +681,7 @@ export function MarkdownRenderer(props) {
animation: 'spin 1s linear infinite',
}}
/>
正在渲染...
{t('正在渲染...')}
</div>
) : (
<MarkdownContent


+ 1
- 1
web/src/components/common/ui/JSONEditor.jsx View File

@@ -661,7 +661,7 @@ const JSONEditor = ({
{hasJsonError && (
<Banner
type='danger'
description={`JSON 格式错误: ${jsonError}`}
description={`${t('JSON 格式错误')}: ${jsonError}`}
className='mb-3'
/>
)}


+ 2
- 0
web/src/components/layout/Footer.jsx View File

@@ -52,6 +52,8 @@ const FooterBar = () => {
<img
src={logo}
alt={systemName}
referrerPolicy='no-referrer'
crossOrigin='anonymous'
className='w-16 h-16 rounded-full bg-gray-800 p-1.5 object-contain'
/>
</div>


+ 5
- 4
web/src/components/layout/PageLayout.jsx View File

@@ -91,6 +91,11 @@ const PageLayout = () => {
if (success) {
statusDispatch({ type: 'set', payload: 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 {
showError('Unable to connect to server');
}
@@ -113,10 +118,6 @@ const PageLayout = () => {
linkElement.href = logo;
}
}
const savedLang = localStorage.getItem('i18nextLng');
if (savedLang) {
i18n.changeLanguage(savedLang);
}
}, [i18n]);

return (


+ 2
- 0
web/src/components/layout/headerbar/HeaderLogo.jsx View File

@@ -44,6 +44,8 @@ const HeaderLogo = ({
<img
src={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'}`}
/>
</div>


+ 0
- 30
web/src/components/layout/headerbar/LanguageSelector.jsx View File

@@ -27,7 +27,6 @@ const LanguageSelector = ({ currentLang, onLanguageChange, t }) => {
position='bottomRight'
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'>
{/* Language sorting: Order by English name (Chinese, English, French, Japanese, Russian) */}
<Dropdown.Item
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'}`}
@@ -35,40 +34,11 @@ const LanguageSelector = ({ currentLang, onLanguageChange, t }) => {
简体中文
</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')}
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
</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>
}
>


+ 1
- 1
web/src/components/playground/CodeViewer.jsx View File

@@ -201,7 +201,7 @@ const CodeViewer = ({ content, title, language = 'json' }) => {
}
return (
formattedContent.substring(0, PERFORMANCE_CONFIG.PREVIEW_LENGTH) +
'\n\n// ... 内容被截断以提升性能 ...'
'\n\n// ... ' + t('内容被截断以提升性能') + ' ...'
);
}, [formattedContent, contentMetrics.isLarge, isExpanded]);



+ 1
- 1
web/src/components/playground/DebugPanel.jsx View File

@@ -146,7 +146,7 @@ const DebugPanel = ({
{t('预览请求体')}
{customRequestMode && (
<span className='px-1.5 py-0.5 text-xs bg-orange-100 text-orange-600 rounded-full'>
自定义
{t('自定义')}
</span>
)}
</div>


+ 2
- 2
web/src/components/playground/MessageContent.jsx View File

@@ -272,7 +272,7 @@ const MessageContent = ({
<div key={index} className='max-w-sm'>
<img
src={imgItem.image_url.url}
alt={`用户上传的图片 ${index + 1}`}
alt={t('用户上传的图片', { index: index + 1 })}
className='rounded-lg max-w-full h-auto shadow-sm border'
style={{ maxHeight: '300px' }}
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'
style={{ display: 'none' }}
>
图片加载失败: {imgItem.image_url.url}
{t('图片加载失败')}: {imgItem.image_url.url}
</div>
</div>
))}


+ 2
- 1
web/src/components/playground/OptimizedComponents.js View File

@@ -74,7 +74,8 @@ export const OptimizedSettingsPanel = React.memo(
prevProps.showSettings === nextProps.showSettings &&
JSON.stringify(prevProps.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 View File

@@ -45,6 +45,7 @@ const SettingsPanel = ({
onCustomRequestBodyChange,
previewPayload,
messages,
channels = [],
}) => {
const { t } = useTranslation();

@@ -176,6 +177,38 @@ const SettingsPanel = ({
/>
</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输入 */}
<div className={customRequestMode ? 'opacity-50' : ''}>
<ImageUrlInput


+ 2
- 2
web/src/components/playground/ThinkingContent.jsx View File

@@ -105,7 +105,7 @@ const ThinkingContent = ({
style={{ color: 'white' }}
className='text-xs mt-0.5 opacity-80 hidden sm:block'
>
来源: {thinkingSource}
{t('来源')}: {thinkingSource}
</Typography.Text>
)}
</div>
@@ -122,7 +122,7 @@ const ThinkingContent = ({
style={{ color: 'white' }}
className='text-xs sm:text-sm font-medium opacity-90'
>
思考中
{t('思考中')}
</Typography.Text>
</div>
)}


+ 5
- 4
web/src/components/playground/configStorage.js View File

@@ -21,6 +21,7 @@ import {
STORAGE_KEYS,
DEFAULT_CONFIG,
} from '../../constants/playground.constants';
import i18next from 'i18next';

const MESSAGES_STORAGE_KEY = 'playground_messages';

@@ -215,16 +216,16 @@ export const importConfig = (file) => {

resolve(importedConfig);
} else {
reject(new Error('配置文件格式无效'));
reject(new Error(i18next.t('配置文件格式无效')));
}
} 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);
} catch (error) {
reject(new Error('导入配置失败: ' + error.message));
reject(new Error(i18next.t('导入配置失败: ') + error.message));
}
});
};

+ 1
- 1
web/src/components/settings/ModelDeploymentSetting.jsx View File

@@ -60,7 +60,7 @@ const ModelDeploymentSetting = () => {
setLoading(true);
await getOptions();
} catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
console.error(error);
} finally {
setLoading(false);


+ 1
- 1
web/src/components/settings/ModelSetting.jsx View File

@@ -95,7 +95,7 @@ const ModelSetting = () => {
await getOptions();
// showSuccess('刷新成功');
} catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
console.error(error);
} finally {
setLoading(false);


+ 137
- 53
web/src/components/settings/OtherSetting.jsx View File

@@ -21,6 +21,7 @@ import React, { useContext, useEffect, useRef, useState } from 'react';
import {
Banner,
Button,
ButtonGroup,
Col,
Form,
Row,
@@ -34,15 +35,26 @@ import { useTranslation } from 'react-i18next';
import { StatusContext } from '../../context/Status';
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 { t } = useTranslation();
const [editingLang, setEditingLang] = useState('zh');
let [inputs, setInputs] = useState({
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: '',
Logo: '',
Footer: '',
@@ -74,8 +86,10 @@ const OtherSetting = () => {

const [loadingInput, setLoadingInput] = useState({
Notice: false,
[LEGAL_USER_AGREEMENT_KEY]: false,
[LEGAL_PRIVACY_POLICY_KEY]: false,
userAgreement: false,
privacyPolicy: false,
termsOfService: false,
usagePolicy: false,
SystemName: false,
Logo: false,
HomePageContent: false,
@@ -88,6 +102,13 @@ const OtherSetting = () => {
setInputs((inputs) => ({ ...inputs, [name]: value }));
};

// 语言切换时同步 form values,确保重新挂载的 TextArea 拿到正确内容
useEffect(() => {
if (formAPISettingGeneral.current) {
formAPISettingGeneral.current.setValues(inputs);
}
}, [editingLang]);

// 通用设置
const formAPISettingGeneral = useRef();
// 通用设置 - Notice
@@ -103,48 +124,19 @@ const OtherSetting = () => {
setLoadingInput((loadingInput) => ({ ...loadingInput, Notice: false }));
}
};
// 通用设置 - UserAgreement
const submitUserAgreement = async () => {
// 通用法律文档保存(同时保存中英文)
const submitLegalDoc = async (docKey, successMsg, errorMsg) => {
const keys = LEGAL_KEYS[docKey];
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) {
console.error(t('用户协议更新失败'), error);
showError(t('用户协议更新失败'));
console.error(t(errorMsg), error);
showError(t(errorMsg));
} 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('设置公告')}
</Button>
<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(
'在此输入用户协议内容,支持 Markdown & HTML 代码',
)}
field={LEGAL_USER_AGREEMENT_KEY}
field={LEGAL_KEYS.userAgreement[editingLang]}
key={`ua_${editingLang}`}
onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }}
@@ -389,17 +396,32 @@ const OtherSetting = () => {
)}
/>
<Button
onClick={submitUserAgreement}
loading={loadingInput[LEGAL_USER_AGREEMENT_KEY]}
onClick={() => submitLegalDoc('userAgreement', '用户协议已更新', '用户协议更新失败')}
loading={loadingInput['userAgreement']}
>
{t('设置用户协议')}
</Button>
<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(
'在此输入隐私政策内容,支持 Markdown & HTML 代码',
)}
field={LEGAL_PRIVACY_POLICY_KEY}
field={LEGAL_KEYS.privacyPolicy[editingLang]}
key={`pp_${editingLang}`}
onChange={handleInputChange}
style={{ fontFamily: 'JetBrains Mono, Consolas' }}
autosize={{ minRows: 6, maxRows: 12 }}
@@ -408,11 +430,73 @@ const OtherSetting = () => {
)}
/>
<Button
onClick={submitPrivacyPolicy}
loading={loadingInput[LEGAL_PRIVACY_POLICY_KEY]}
onClick={() => submitLegalDoc('privacyPolicy', '隐私政策已更新', '隐私政策更新失败')}
loading={loadingInput['privacyPolicy']}
>
{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.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>
</Card>
</Form>


+ 6
- 0
web/src/components/settings/PaymentSetting.jsx View File

@@ -24,6 +24,7 @@ import SettingsPaymentGateway from '../../pages/Setting/Payment/SettingsPaymentG
import SettingsPaymentGatewayStripe from '../../pages/Setting/Payment/SettingsPaymentGatewayStripe';
import SettingsPaymentGatewayCreem from '../../pages/Setting/Payment/SettingsPaymentGatewayCreem';
import SettingsPaymentGatewayWechat from '../../pages/Setting/Payment/SettingsPaymentGatewayWechat';
import SettingsPaymentGatewayAlipay from '../../pages/Setting/Payment/SettingsPaymentGatewayAlipay';
import { API, showError, toBoolean } from '../../helpers';
import { useTranslation } from 'react-i18next';

@@ -101,6 +102,8 @@ const PaymentSetting = () => {
case 'StripeMinTopUp':
case 'WechatPayUnitPrice':
case 'WechatPayMinTopUp':
case 'AlipayUnitPrice':
case 'AlipayMinTopUp':
newInputs[item.key] = parseFloat(item.value);
break;
default:
@@ -152,6 +155,9 @@ const PaymentSetting = () => {
<Card style={{ marginTop: '10px' }}>
<SettingsPaymentGatewayWechat options={inputs} refresh={onRefresh} />
</Card>
<Card style={{ marginTop: '10px' }}>
<SettingsPaymentGatewayAlipay options={inputs} refresh={onRefresh} />
</Card>
</Spin>
</>
);


+ 1
- 1
web/src/components/settings/RateLimitSetting.jsx View File

@@ -64,7 +64,7 @@ const RateLimitSetting = () => {
await getOptions();
// showSuccess('刷新成功');
} catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
} finally {
setLoading(false);
}


+ 1
- 1
web/src/components/settings/RatioSetting.jsx View File

@@ -83,7 +83,7 @@ const RatioSetting = () => {
setLoading(true);
await getOptions();
} catch (error) {
showError('刷新失败');
showError(t('刷新失败'));
} finally {
setLoading(false);
}


+ 40
- 1
web/src/components/settings/SystemSetting.jsx View File

@@ -78,6 +78,7 @@ const SystemSetting = () => {
WeChatServerToken: '',
WeChatAccountQRCodeImageURL: '',
TurnstileCheckEnabled: '',
CaptchaEnabled: '',
TurnstileSiteKey: '',
TurnstileSecretKey: '',
RegisterEnabled: '',
@@ -100,6 +101,7 @@ const SystemSetting = () => {
LinuxDOClientSecret: '',
LinuxDOMinimumTrustLevel: '',
ServerAddress: '',
DefaultLanguage: '',
// SSRF防护配置
'fetch_setting.enable_ssrf_protection': true,
'fetch_setting.allow_private_ip': '',
@@ -179,6 +181,7 @@ const SystemSetting = () => {
case 'TelegramOAuthEnabled':
case 'RegisterEnabled':
case 'TurnstileCheckEnabled':
case 'CaptchaEnabled':
case 'EmailDomainRestrictionEnabled':
case 'EmailAliasRestrictionEnabled':
case 'SMTPSSLEnabled':
@@ -317,6 +320,10 @@ const SystemSetting = () => {
await updateOptions([{ key: 'ServerAddress', value: ServerAddress }]);
};

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

const submitSMTP = async () => {
const options = [];

@@ -716,7 +723,7 @@ const SystemSetting = () => {
<Row
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
field='ServerAddress'
label={t('服务器地址')}
@@ -726,10 +733,33 @@ const SystemSetting = () => {
)}
/>
</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>
<Button onClick={submitServerAddress}>
{t('更新服务器地址')}
</Button>
<Button onClick={submitDefaultLanguage}>
{t('保存默认语言')}
</Button>
</Form.Section>
</Card>

@@ -1033,6 +1063,15 @@ const SystemSetting = () => {
>
{t('允许 Turnstile 用户校验')}
</Form.Checkbox>
<Form.Checkbox
field='CaptchaEnabled'
noLabel
onChange={(e) =>
handleCheckboxChange('CaptchaEnabled', e)
}
>
{t('图片验证码')}
</Form.Checkbox>
</Col>
<Col xs={24} sm={24} md={12} lg={12} xl={12}>
<Form.Checkbox


+ 1
- 1
web/src/components/settings/personal/cards/NotificationSettings.jsx View File

@@ -533,7 +533,7 @@ const NotificationSettings = ({
<CodeViewer
content={{
type: 'quota_exceed',
title: '额度预警通知',
title: t('额度预警通知'),
content:
'您的额度即将用尽,当前剩余额度为 {{value}}',
values: ['$0.99'],


Some files were not shown because too many files changed in this diff

Loading…
Cancel
Save