| Author | SHA1 | Message | Date |
|---|---|---|---|
|
|
caf457ce9e |
feat: 错误时记录上游响应体,支持流式
- 提取 TruncateBody 公共函数,截断到 2KB 避免日志过大 - 新增 handleResponsesStreamError 统一 SSE 错误提取逻辑 - 流式/非流式 Responses API 错误路径均捕获 UpstreamBody - 添加 upstream_body 和 truncate_body 单元测试 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
575f84064e |
feat: codex API key 模式 + 前端凭证回显修复
- codex adaptor 支持 API key 和 OAuth 两种认证模式 - 提取 setupOAuthHeader 和 shouldUseChatCompletionsViaResponses - 修复编辑页 codex_credential_mode 切换/回显不同步问题 - ResponsesStreamResponse 添加 Error 字段支持独立错误事件 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
ca91922904 | docs: add log chat/upstream id design | 3 months ago |
|
|
332d5a62f2 | feat: support codex api key credentials | 3 months ago |
|
|
1c0898257c | feat: add redemption remarks | 3 months ago |
|
|
f18e51e3a9 | feat: add redemption remark support | 3 months ago |
|
|
fae24fdf37 | test: stabilize channel affinity usage cache tests | 3 months ago |
|
|
fc0257f6ca | chore: ignore local worktrees | 3 months ago |
|
|
d8272e7707 |
feat(channel): 添加渠道"对外名称"(public_name)字段
为 Channel 模型新增 public_name 字段,让管理员可以为每个渠道 设置用户可见的友好名称(如"标准通道"、"高速通道"),替代前端 硬编码的"通道一/二/三"。Playground 通道选择器展示对外名称, "渠道"统一改为"通道"。新建渠道时对外名称必填,编辑时可选。 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
73d10b5799 |
feat: 错误日志记录上游 request-id 和响应体,Playground 渠道路由改为 header 传递
- 错误日志新增 upstream_request_id(从 Anthropic/OpenAI 响应 header 提取)和 upstream_body(截断 2KB) - 修复 RelayErrorHandler 内部 WithOpenAIError/NewOpenAIError 分支丢失上游字段的 bug - Playground 渠道和分组改为通过 X-Channel-Id/X-Group header 传递,而非 body 字段 - Distributor 中间件支持从 header 回退读取 channel_id 和 group Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
5b8303ae81 |
fix: 登录后跳转来源页,充值金额校验优化
- LoginForm: 登录成功后跳转回登录前的页面,而非固定 /console - RechargeCard: 输入金额低于最低值时显示警告提示 - .dockerignore: 排除 .claude、.plans、脚本等非项目文件 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
8bb2883d7e |
feat(home): 首页定价卡片添加"去体验"按钮,跳转 Playground
- PricingCardView 新增 showTryButton 属性,点击跳转 Playground 并预选模型 - 移除 PricingCardView 中未使用的 props (selectedGroup, currency 等) - 首页页脚路由重命名: terms→user-agreement, usage-policy→privacy-policy - i18n: "操练场"更名为"对话" Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
a0ac37b2a9 |
fix: 替换 println 为 SysLog,修复 logger nil context 崩溃
- logger: 增加 ctx nil 检查,避免系统级日志 panic - relay: 将 DebugEnabled 下的 println 替换为 SysLog - price: 修复渠道定价在 ChannelMeta 未初始化时的获取失败问题 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
74cc8c0d56 |
feat(playground): 添加渠道选择功能,支持指定渠道体验模型
- 新增 /api/user/model_channels 接口,返回模型可用渠道及默认渠道 - Playground 设置面板添加渠道选择下拉框 - Distribute 中间件支持从请求体读取 channel_id 指定渠道 - 支持 URL 参数 ?model=xxx 直接选择模型 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
a34d44a824 |
fix(user): 修复从节点同步用户余额显示为 0 的问题
同步用户(synced)的实际余额存储在 synced_quota 字段, 但用户列表/详情接口返回的是 quota 字段(值为 0)。 新增 ApplySyncedQuota 方法,在 API 返回前将 synced_quota 赋值给 quota,使前端能正确显示余额。 同时优化 GetUserModels:map 去重替代 O(n) 线性扫描, 禁用模型查询移至 model 层,新增 StringsSubtract 工具函数。 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
967309fb56 |
feat: 添加邮箱后缀注册额度规则功能
管理员可在用户管理页面配置「邮箱后缀→初始额度」映射规则, 用户注册时根据邮箱后缀自动匹配并发放对应额度(替代默认额度)。 - 新增 email_quota_rule 数据表 + CRUD API(管理员权限) - 内存缓存匹配,启动时加载,增删改时刷新 - 注册流程 Insert/InsertWithTx/FinalizeOAuthUserCreation 同步支持 - 前端用户管理页新增 Tab 展示规则管理卡片 - 含 24 个测试(model 16 + controller 8) Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
83f767b40b |
Merge branch 'worktree-feat-captcha'
# Conflicts: # web/src/components/topup/RechargeCard.jsx |
3 months ago |
|
|
d24d68a916 |
fix(i18n): 修复令牌和充值页面货币显示不一致问题
- 令牌创建:快捷选项标签根据系统货币设置动态生成,替换硬编码美元 - 充值页面:输入框显示值与预设卡片统一换算为本地货币,内部值仍为原始单位 - 使用 getQuotaPerUnit() 替代魔数 500000 - 消除 RechargeCard 中 getCurrencyConfig() 的 N+1 重复调用 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
fe88bc7758 | Merge branch 'worktree-feat-captcha' | 3 months ago |
|
|
1912dd3b72 |
feat: 添加图片验证码功能,防止脚本批量注册
- 后端:新增 GET /api/captcha 接口,使用 base64Captcha 生成图片验证码 - 后端:发送邮箱验证码时校验图片验证码(CaptchaEnabled 开关控制) - 前端:注册表单邮箱验证码前增加图片验证码输入,支持点击刷新 - 管理后台:系统设置新增「图片验证码」开关 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
20163dbad5 |
fix(legal): 法律文档页面切换语言时强制重新加载内容
给 DocumentRenderer 添加 key={lang} 属性,确保切换语言时组件
重新挂载并从 API 获取对应语言的内容,而非使用缓存内容。
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
3 months ago |
|
|
56e8d14485 |
fix: 语言切换后 TextArea 内容丢失
添加 useEffect 在 editingLang 变化时同步 form values, 确保重新挂载的 TextArea 能从 React state 中读取正确内容。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
65bd1dabb8 |
fix: 法律文档编辑器切换语言时 TextArea 内容未刷新
给每个 TextArea 添加 key={editingLang} 强制组件在语言切换时重新挂载。
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
3 months ago |
|
|
934b4b10a7 |
feat: 法律文档双语配置(中/英)+ 精简语言支持至中英双语
- LegalSettings 字段拆分为 _zh/_en 双语对 - API 支持 ?lang= 参数返回对应语言内容 - 后台设置页每个文档配语言切换按钮(中文/English) - 前端页面根据 UI 语言自动请求对应内容 - 移除 fr/ru/ja/vi/zh-TW 语言支持 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
676902ed58 |
Merge branch 'worktree-feat-terms-usage-policy'
# Conflicts: # web/src/i18n/locales/en.json # web/src/i18n/locales/fr.json # web/src/i18n/locales/ja.json # web/src/i18n/locales/ru.json # web/src/i18n/locales/vi.json # web/src/i18n/locales/zh-CN.json # web/src/i18n/locales/zh-TW.json |
3 months ago |
|
|
2757f77011 |
feat(i18n): 补全服务条款和使用政策的多语言翻译
在 zh-CN、en、zh-TW、ja、fr、ru、vi 七个语言文件中添加 服务条款和使用政策的完整翻译 key。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
65c6446a30 | Merge branch 'worktree-feat-terms-usage-policy' | 3 months ago |
|
|
1cc75aa9a2 |
feat: 服务条款和使用政策页面,支持后台 Markdown 配置
复用现有 LegalSettings 模式,新增 TermsOfService 和 UsagePolicy 字段, 添加 /terms 和 /usage-policy 前端页面及 API 端点。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
ed195cada5 |
fix(frontend): sort_order 默认值显示为"未设置",编辑弹窗增加提示
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
d5157a779c |
fix(pricing): 统一 sort_order 排序逻辑,有 meta 但未设置的不优先
之前有 meta 记录但 sort_order=999999 的模型排在没 meta 记录的模型前面, 导致部分"未设置排序"的模型仍然挤在前面。现在统一处理:不论是否有 meta 记录,sort_order=999999 的都视为未设置,排在后面。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
54e53108ec |
fix(sort): sort_order 默认值改为 999999,简化排序逻辑
未设置排序的记录 sort_order=999999 自然排在后面,无需 CASE WHEN。 - GORM 默认值 default:0 → default:999999 - 数据库迁移:将现有 sort_order=0 的记录更新为 999999 - 回退 CASE WHEN 排序逻辑,恢复简单的 sort_order ASC, id ASC - 前端编辑弹窗默认值同步改为 999999 - 表格列中 999999 显示为空(表示未设置) Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
3b99c83e32 |
fix(sort): sort_order=0 的记录排到最后,非0值按升序排列
模型和供应商查询统一排序规则:sort_order=0 视为未设置排到最后, 非0值按升序排列。涉及 GetAllModels、SearchModels、GetAllVendors、 SearchVendors 以及 pricing 的内存排序。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
518f1ab87f |
refactor(frontend): 移除拖拽排序,改为 sort_order 数值输入
回退 @dnd-kit 拖拽实现,改为在编辑弹窗中直接设置 sort_order 数值: - 模型编辑弹窗新增「排序」InputNumber 字段 - 供应商编辑弹窗新增「排序」InputNumber 字段 - 模型表格新增 sort_order 显示列 - 移除 @dnd-kit 依赖及所有 DnD 相关代码 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
65880fc69b |
feat(frontend): Vendor Tab 和 Model 表格支持拖拽排序,移除定价页字母排序
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
5e0f38d26c |
chore(frontend): 安装 @dnd-kit 拖拽排序库
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
e507897f21 |
feat(pricing): updatePricing 按 sort_order 排序 vendors 和 models
- vendorsList 按 vendor.SortOrder 升序排列,相同则按 ID 排序 - pricingMap 按 model.SortOrder 升序排列,相同则按模型名字典序 - 修复默认通道设置的 API 错误响应未正确展示的问题 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
749064932c |
feat: 新增 PUT /api/vendors/reorder 和 /api/models/reorder 批量排序接口
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
f8859b7c35 |
feat: Vendor 和 Model 添加 sort_order 字段,查询按 sort_order ASC 排序
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
3a29f3772f |
feat(channel): 添加模型默认通道功能,支持优先路由和卡片标识
- 新增 is_default 字段和缓存层,支持管理员为模型指定默认通道 - Distribute 中间件优先级调整:Token 指定 → 默认通道 → 亲和性 → 随机 - 模型定价卡片和详情弹窗展示默认通道 amber 标识 - 管理后台定价页面新增星标切换默认通道 - 新增 set_default / clear_default API 和 6 个单元测试 - 简化卡片价格显示(移除内联缓存价格,改为详情弹窗展示) Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
93e331b624 |
feat: 支持 LOGO_FILE_PATH 环境变量指定本地 Logo,优化定价与语言设置
- 新增 LOGO_FILE_PATH 环境变量,优先级高于数据库配置,支持本地文件服务 - 渠道定价高级字段(缓存/图片/音频)不再回退全局默认值,未设置直接返回 0 - 缓存价格单位从表头移到具体价格值,新增渠道 ID 复制功能 - 修复默认语言在用户已有偏好时仍被覆盖的问题 - img 标签统一添加 referrerPolicy/crossOrigin 防止跨域问题 - 新增缓存倍率和最小余额阈值的详细说明文本 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
94f24b29f7 |
style: 首页文案"全球"改为"顶级"
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
f0f0215b20 |
fix(pricing): 修复缓存倍率写入默认值1和缓存创建token丢失
1. CopyGlobalPricing 检查 ratio 查找的 bool 返回值,未找到时写 0 而非写入 fallback 值 1,避免覆盖全局正确的 0.1 2. Claude 响应使用 GetCacheCreationTotalTokens() 替代直接读 CacheCreationInputTokens,兼容新版本子对象格式 3. 新增 GetAudioRatioV2/GetAudioCompletionRatioV2 带 bool 返回 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
3 months ago |
|
|
387f6c1ae0 |
feat(pricing): 定价数据源切换到渠道表,新增缓存价格展示
- 后端 pricing API 从 channel_pricings 表获取实际定价,选取最便宜渠道 - 提取 applyGlobalDefault 辅助函数消除全局回退逻辑重复 - price.go 重构扩展比率为局部变量,简化回退逻辑 - 定价卡片新增缓存读取/创建价格,改为两行布局防止溢出 - ChannelPricingCard 缓存列显示实际价格而非倍率 - 修复移动端 hero 区域 padding 过大 - 默认标签页标题改为 Loading... - 新增缓存读取/创建 i18n 翻译(7 语言) Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
00058cd671 |
refactor(channel-pricing): 提取辅助方法消除重复代码,修复缓存一致性
- 提取 PriceData.ApplyChannelPricingRatios 消除 ModelPriceHelper 和 UpdatePriceDataForChannelPricing 中重复的 ~30 行比率回退逻辑 - 提取 ChannelPricing.ApplyFields 消除 controller 中 4 处相同的字段赋值 - 提取 setCache/removeCache 辅助函数统一写穿透缓存操作 - 将 claudeCacheCreation1hMultiplier 常量移至 types 包避免循环依赖 - 修复 BatchUpsertChannelPricing 成功后未刷新内存缓存的 bug Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
42082311d3 |
fix: 修复默认语言不生效的问题
移除 PageLayout useEffect 中多余的 localStorage 语言恢复逻辑。 i18next-browser-languagedetector 在初始化时已自动处理 localStorage, 同步的 changeLanguage 调用会覆盖异步 loadStatus 设置的管理员默认语言。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
9eb58c684e |
feat: 后台设置默认语言
管理员可在系统设置页面配置全局默认语言,未登录用户和未设置 语言偏好的已登录用户将强制使用该语言,忽略浏览器语言检测。 - common/constants.go: 新增 DefaultLanguage 变量 - model/option.go: 新增 DefaultLanguage option handler - controller/misc.go: /api/status 暴露 default_language - SystemSetting.jsx: 添加语言下拉框到通用设置 - PageLayout.jsx: 未登录用户应用默认语言 - UserContext.jsx: 无偏好用户回退到默认语言 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
a2b27ea2a6 |
docs: 后台设置默认语言实施计划
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
436fcb405b |
docs: 后台设置默认语言功能设计文档
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
8d22c4e80e |
style(home): 临时隐藏首页工具链/核心价值/工作流/生态伙伴 section
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
355757404b | merge: feat/channel-pricing-extended → master | 3 months ago |
|
|
235f7c6e5f |
refactor(channel-pricing): 提取 ParseTagIds 辅助函数 + 移除调试 console.log
- 将 tag ID 解析逻辑提取为 model.ParseTagIds,消除 model 层和 controller 层重复代码
- 统一 TrimSpace 处理(之前 controller 版本漏了)
- 移除 ChannelPricingView 残留的 console.log('Expanded keys:')
- 删除多余空行
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
|
3 months ago |
|
|
f3577590bc |
test(channel-pricing): 添加 API 端到端集成测试脚本
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
129e17ac42 |
feat(channel-pricing): 前端支持缓存/图片/音频倍率编辑和展示
- ChannelPricingView.jsx: initValues 添加 5 个高级比例字段 + 按量计费模式下新增高级比例表单区块 - ChannelPricingCard.jsx: tableData 映射 cache_ratio/cache_creation_ratio + 条件列展示 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
c84c26b5a2 |
feat(channel-pricing): GetChannelPricingByModelWithChannelInfo 支持 CASE WHEN 回退 + 扩展字段
在 SQL 查询中使用 CASE WHEN > 0 回退策略,让渠道定价的扩展比率 (cache_ratio, cache_creation_ratio, image_ratio, audio_ratio, audio_completion_ratio) 在未设置时自动回退到全局默认值, 而不是像 COALESCE 那样被零值拦截。同时更新 GROUP BY 子句包含所有 新选择的列,确保 MySQL 兼容。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
d5d4714908 |
test(channel-pricing): 缓存写穿 + 字段默认值单元测试
验证 ChannelPricing 的 Insert/Update/Delete 操作正确更新内存缓存, 以及未设置的扩展字段(CacheRatio 等)默认为零值。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
646f37dd4b |
feat(channel-pricing): Controller 扩展 API 支持新字段 + 输入校验 + 操作日志
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
3a24fd598f |
feat(channel-pricing): ModelPriceHelper + UpdatePriceDataForChannelPricing 适配新签名,支持扩展比率覆盖
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
e5f029f91e |
feat(channel-pricing): InitDB 启动时全量加载渠道定价缓存
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
2d4c73d3aa |
refactor(channel-pricing): 结构体新增扩展字段 + 缓存重写为全量加载写穿
- ChannelPricing 结构体新增 5 个扩展计费字段: cache_ratio, cache_creation_ratio, image_ratio, audio_ratio, audio_completion_ratio - 缓存机制从 TTL+惰性加载重写为全量加载+写穿模式,消除 DB 查询延迟 - Insert/Update/Delete 改为写穿缓存(直接更新内存,不再整表失效) - 新增 LoadChannelPricingCache 启动时全量加载函数 - 删除 RefreshChannelPricingCache/InvalidateChannelPricingCache(不再需要) - BatchUpsertChannelPricing DoUpdates 列表同步新增 5 字段 - GetEffectivePricing 签名改为 (*ChannelPricing, bool)(调用方适配在后续 Task) Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
|
|
156618fdad |
merge: feat/alipay-payment → master
支付宝当面付扫码支付集成 + 代码重构 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
b514a0798b |
refactor(payment): 合并微信/支付宝支付重复代码,删除冗余文件
- 提取通用 calcPayMoney/calcMinTopup 函数,消除微信和支付宝控制器中的重复计费逻辑 - 合并 RechargeAlipay/RechargeWechat 为 rechargeByQRCodePayment 内部函数,删除 model/topup_alipay.go - 复用已有的 wrapAsPEM 替换支付宝专用的 wrapAlipayPublicKey - 删除被 QRCodePayModal 替代的 WechatPayQRCodeModal.jsx - 修复 controller/topup.go 中支付宝代码块的缩进错误 净减 264 行代码。 Co-Authored-By: Claude <noreply@anthropic.com> |
3 months ago |
|
|
e8ade3abae |
feat(payment): 集成支付宝当面付扫码支付 + 补全前端 i18n 硬编码中文
支付宝支付: - 后端:新增 topup_alipay.go、payment_alipay.go,实现当面付(扫码支付)下单、 回调验签、订单状态查询完整流程,支持 RSA2 签名 - 前端:新增 QRCodePayModal 通用二维码支付弹窗(泛化微信支付弹窗)、 SettingsPaymentGatewayAlipay 管理配置页面、充值流程接入支付宝 - 充值历史:管理员视图新增用户邮箱列,后端 fillTopUpEmails 批量填充 i18n 补全: - 54 个缺失翻译 key 添加到全部 7 个 locale 文件(zh-CN/zh-TW/en/fr/ja/ru/vi) - 涵盖渠道名称、仪表盘标签、兑换码状态、控制台时间筛选、Playground 错误消息、 Dashboard 设置页面提示等 - 14 个源码文件中约 31 处 showError/showSuccess/showWarning 硬编码中文改用 t() - ChannelsColumnDefs、helpers、services 中散落的硬编码中文统一国际化 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
3 months ago |
| @@ -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 | |||
| @@ -36,3 +36,4 @@ | |||
| # ============================================ | |||
| # Mark web frontend as vendored so GitHub recognizes this as a Go project | |||
| electron/** linguist-vendored | |||
| .dockerignore text eol=lf | |||
| @@ -23,6 +23,7 @@ plans | |||
| docs/plans/ | |||
| CLAUDE.md | |||
| .claude | |||
| .worktrees/ | |||
| logs/ | |||
| docs/superpowers | |||
| @@ -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) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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 // 是否启用邮箱域名限制 | |||
| @@ -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)) | |||
| @@ -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, | |||
| }) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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(), "当前凭证方式不支持查看用量") | |||
| } | |||
| @@ -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()) | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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{ | |||
| @@ -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, | |||
| }, | |||
| }) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| @@ -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") | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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": "", | |||
| @@ -189,6 +189,7 @@ | |||
| | key | string | 兑换码(32字符,唯一) | | |||
| | status | int | 状态:1=启用,2=已使用,3=已禁用 | | |||
| | name | string | 兑换码名称 | | |||
| | remark | string | 备注 | | |||
| | quota | int | 额度值 | | |||
| | created_time | int64 | 创建时间 | | |||
| | redeemed_time | int64 | 兑换时间 | | |||
| @@ -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) | |||
| @@ -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 表) | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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= | |||
| @@ -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) | |||
| @@ -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)) | |||
| @@ -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 != "" { | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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")) | |||
| } | |||
| @@ -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)) | |||
| } | |||
| } | |||
| @@ -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 | |||
| }) | |||
| } | |||
| @@ -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": | |||
| @@ -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) | |||
| } | |||
| @@ -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, "应选最便宜的按次渠道") | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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,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 | |||
| @@ -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,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 | |||
| }) | |||
| } | |||
| @@ -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 { | |||
| @@ -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() | |||
| } | |||
| @@ -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") | |||
| } | |||
| @@ -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: | |||
| @@ -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() | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| }) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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 { | |||
| @@ -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)) | |||
| } | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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()) | |||
| @@ -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,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 | |||
| } | |||
| @@ -93,7 +93,7 @@ func runCodexCredentialAutoRefreshOnce() { | |||
| continue | |||
| } | |||
| oauthKey, err := parseCodexOAuthKey(rawKey) | |||
| oauthKey, err := common.ParseCodexOAuthCredential(rawKey) | |||
| if err != nil { | |||
| continue | |||
| } | |||
| @@ -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()) | |||
| } | |||
| @@ -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) | |||
| }) | |||
| } | |||
| @@ -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() | |||
| @@ -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) | |||
| @@ -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() { | |||
| @@ -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 "=== 集成测试完成 ===" | |||
| @@ -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. | |||
| @@ -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) | |||
| @@ -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> | |||
| @@ -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={ | |||
| @@ -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'> | |||
| @@ -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> | |||
| @@ -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> | |||
| @@ -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'> | |||
| @@ -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 | |||
| @@ -661,7 +661,7 @@ const JSONEditor = ({ | |||
| {hasJsonError && ( | |||
| <Banner | |||
| type='danger' | |||
| description={`JSON 格式错误: ${jsonError}`} | |||
| description={`${t('JSON 格式错误')}: ${jsonError}`} | |||
| className='mb-3' | |||
| /> | |||
| )} | |||
| @@ -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> | |||
| @@ -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 ( | |||
| @@ -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> | |||
| @@ -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> | |||
| } | |||
| > | |||
| @@ -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]); | |||
| @@ -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> | |||
| @@ -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> | |||
| ))} | |||
| @@ -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) | |||
| ); | |||
| }, | |||
| ); | |||
| @@ -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 | |||
| @@ -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> | |||
| )} | |||
| @@ -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)); | |||
| } | |||
| }); | |||
| }; | |||
| @@ -60,7 +60,7 @@ const ModelDeploymentSetting = () => { | |||
| setLoading(true); | |||
| await getOptions(); | |||
| } catch (error) { | |||
| showError('刷新失败'); | |||
| showError(t('刷新失败')); | |||
| console.error(error); | |||
| } finally { | |||
| setLoading(false); | |||
| @@ -95,7 +95,7 @@ const ModelSetting = () => { | |||
| await getOptions(); | |||
| // showSuccess('刷新成功'); | |||
| } catch (error) { | |||
| showError('刷新失败'); | |||
| showError(t('刷新失败')); | |||
| console.error(error); | |||
| } finally { | |||
| setLoading(false); | |||
| @@ -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> | |||
| @@ -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> | |||
| </> | |||
| ); | |||
| @@ -64,7 +64,7 @@ const RateLimitSetting = () => { | |||
| await getOptions(); | |||
| // showSuccess('刷新成功'); | |||
| } catch (error) { | |||
| showError('刷新失败'); | |||
| showError(t('刷新失败')); | |||
| } finally { | |||
| setLoading(false); | |||
| } | |||
| @@ -83,7 +83,7 @@ const RatioSetting = () => { | |||
| setLoading(true); | |||
| await getOptions(); | |||
| } catch (error) { | |||
| showError('刷新失败'); | |||
| showError(t('刷新失败')); | |||
| } finally { | |||
| setLoading(false); | |||
| } | |||
| @@ -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 | |||
| @@ -533,7 +533,7 @@ const NotificationSettings = ({ | |||
| <CodeViewer | |||
| content={{ | |||
| type: 'quota_exceed', | |||
| title: '额度预警通知', | |||
| title: t('额度预警通知'), | |||
| content: | |||
| '您的额度即将用尽,当前剩余额度为 {{value}}', | |||
| values: ['$0.99'], | |||