Compare commits

...

229 Commits

Author SHA1 Message Date
  fengsilin 4533c24fa1 merge: dynamic video asset channel selection 3 days ago
  fengsilin f6321b3910 feat(channel): support dynamic video asset channel selection 3 days ago
  fengsilin 3fbf48313a feat(pricing): support usage billing and endpoint fixes 3 days ago
  fengsilin dd14839996 merge: tianyiyun seedance channel 3 days ago
  fengsilin 678add297f feat(channel): add tianyiyun seedance channel 3 days ago
  fengsilin f100f9e409 docs: design tianyiyun seedance channel 3 days ago
  fengsilin ffb3991654 Merge branch 'feat/kling-aiping' 6 days ago
  fengsilin 49ca8a6ae7 fix(kling): 修复原生路由和展示价格保存 6 days ago
  fengsilin 6f805a6177 refactor: update aiping doubao upstream API paths to multimodal/sd endpoints 1 week ago
  fengsilin 56400b6c61 feat: KlingAiping 渠道适配 + 展示价格管理后台 1 week ago
  fengsilin a61ca078fc fix: localize display pricing units 1 week ago
  fengsilin fee9504107 feat: render display pricing on pricing page 1 week ago
  fengsilin cc97834e5b feat: add display pricing frontend helpers 1 week ago
  fengsilin ced4d7f323 feat: expose model display pricing 1 week ago
  fengsilin de117ac932 feat: add model display pricing admin api 1 week ago
  fengsilin 18a3495968 feat: add model display pricing settings 1 week ago
  fengsilin 302987989b docs: expand display pricing frontend spec 1 week ago
  fengsilin 6dd8c19ea8 docs: clarify model display pricing design 1 week ago
  fengsilin 6dd9bc6dec docs: add model display pricing design 1 week ago
  fengsilin c9988af52b feat: 用户创建同步到从节点及 i18n 翻译 2 weeks ago
  fengsilin 84d364d76a feat: 使用日志计费详情展开与 completion_tokens 修复 2 weeks ago
  fengsilin 2e40bab478 feat: 模型定价矩阵价格取中位数展示 2 weeks ago
  fengsilin ec7c76f221 merge: feat/video-pricing-table-codex into master 2 weeks ago
  fengsilin 0e3b98efe0 chore: commit relay SSE error handling fixes and E2E mock server 2 weeks ago
  fengsilin 139b6a6075 feat: store relay capture records as json 2 weeks ago
  fengsilin 65551e8003 refactor: prepare relay capture json helpers 2 weeks ago
  fengsilin 19d55dc255 docs: design relay capture json storage 2 weeks ago
  fengsilin 1d01c4cd64 feat: 增强流式响应错误处理,SSE内嵌错误正确曝光为API错误 2 weeks ago
  fengsilin 696c549a66 feat(video): add doubao aiping video pricing support 2 weeks ago
  fengsilin 99168a9660 test: stabilize model pricing e2e runtime 3 weeks ago
  fengsilin aeff58b122 test: stabilize model pricing generator e2e 3 weeks ago
  fengsilin 5451d801f3 test: cover model pricing e2e workflow 3 weeks ago
  fengsilin a4b824952f test: make model pricing tab selector clickable 3 weeks ago
  fengsilin 32014dde92 test: add model pricing e2e selectors 3 weeks ago
  fengsilin eaae02e767 test: ignore e2e artifacts 3 weeks ago
  fengsilin 75285bcabd test: align playwright runtime with chromium 3 weeks ago
  fengsilin b976b486af test: avoid shell by default in e2e process helper 3 weeks ago
  fengsilin 568a91dfcf test: harden e2e helper error handling 3 weeks ago
  fengsilin eb285055ac test: harden e2e process cleanup 3 weeks ago
  fengsilin 32603ba385 docs(test): video pricing table test plan with mock server design 3 weeks ago
  fengsilin b91d07b0da test: add e2e process helpers 3 weeks ago
  fengsilin 9837e80fac test: add e2e scripts 3 weeks ago
  fengsilin cf2c47b826 docs: add e2e test foundation plan 3 weeks ago
  fengsilin 915efefaec docs: add e2e test foundation design 3 weeks ago
  fengsilin 3ac94e2e6f test(task-pricing): cover model pricing api and remix snapshots 3 weeks ago
  fengsilin 19bdae5509 fix(stream): default non-positive streaming timeout 3 weeks ago
  fengsilin 089bc88015 docs(task-pricing): document video pricing table testing 3 weeks ago
  fengsilin 43b0bf293d test(task-pricing): cover video usage billing flows 3 weeks ago
  fengsilin b36edf7f43 feat(ui): add model multidimensional pricing editor 3 weeks ago
  fengsilin aff3284d02 feat(ui): add model pricing config utilities 3 weeks ago
  fengsilin f3778e3931 test(task-pricing): cover doubao usage parsing 3 weeks ago
  fengsilin 285323d765 feat(task-pricing): expose model pricing rules api 3 weeks ago
  fengsilin 3a52f31608 feat(task-pricing): gate channels for matrix usage billing 3 weeks ago
  fengsilin db8eacd558 feat(task-pricing): persist matrix billing snapshots 3 weeks ago
  fengsilin 8f2b014366 feat(task-pricing): add matrix lookup decisions 3 weeks ago
  fengsilin ae1e68792a feat(task-pricing): resolve pricing dimensions from task requests 3 weeks ago
  fengsilin b4301a53ae feat(task-pricing): validate and cache model pricing rules 3 weeks ago
  fengsilin 95060ea0aa feat(task-pricing): add pricing config and decision types 3 weeks ago
  fengsilin 8f89828ca3 fix(user-migration): 禁止冲突用户非法合并 1 month ago
  fengsilin 2f7a71d61e fix(user-migration): 补齐主节点保护与导入校验 1 month ago
  fengsilin 9aaa6baa4c merge: user migration selective batches into master 1 month ago
  fengsilin 65e6a82d34 feat(user-migration): support selective batches and cancellation 1 month ago
  fengsilin ca3f8765bb test(user-migration): complete migration coverage and e2e flow 1 month ago
  fengsilin a72d76522f fix(deploy): 部署脚本 tag 格式与 Makefile 对齐,加入 commit hash 和 dirty 标记 1 month ago
  fengsilin de993c5605 feat(docker): 镜像 tag 加入 commit hash,未提交代码标记 dirty 1 month ago
  fengsilin e9d5d51b37 fix(sync): 修复 Slave 端编辑同步用户时误报余额不可修改的问题 1 month ago
  fengsilin 2544556f5d fix(ratio): 缓存读取倍率为 0 时前端不显示,补充 extraText 说明 1 month ago
  fengsilin dce7631092 docs(ratio): 缓存创建倍率 extraText 补充设置为 0 时隐藏说明 1 month ago
  fengsilin 1c15ee996d feat(ratio): 升级未设置倍率模型编辑弹窗,对齐可视化编辑器 1 month ago
  fengsilin 182ab5ff83 refactor(ratio): move advanced ratios to edit modal in unset models page 1 month ago
  fengsilin 1fdf0a3f47 feat(ratio): add advanced ratio fields to unset models editor 1 month ago
  fengsilin d36a686f6b fix(ratio): preserve 0 as valid value for advanced ratio fields 1 month ago
  fengsilin c824381642 fix(ratio): read advanced ratio values from form on submit 1 month ago
  fengsilin 7b34fbc88a feat(pricing): hide cache creation price when ratio is 0 1 month ago
  fengsilin a8c739e466 feat(ui): add extraText descriptions and fix placeholder for advanced ratios 1 month ago
  fengsilin 192eb70862 feat(save): serialize and save advanced ratio fields 1 month ago
  fengsilin 91e648c530 feat(add): handle advanced ratios in addOrUpdateModel 1 month ago
  fengsilin 194ebbaa56 feat(edit): populate advanced ratios in edit modal 1 month ago
  fengsilin 2ca560ccef feat(ui): add advanced ratios form section 1 month ago
  fengsilin 0ea023037a feat(ratio): parse advanced ratio fields in data initialization 1 month ago
  fengsilin dab78ff2c4 feat(i18n): add advanced ratio translations 1 month ago
  fengsilin a9fc624ff2 docs: add extended visual ratio settings implementation plan 1 month ago
  fengsilin b699dda2ad docs: add extended visual ratio settings design spec 1 month ago
  fengsilin 491e9b8ee4 fix(pricing): GetCompletionRatio 优先使用用户设置值而非硬编码默认值 1 month ago
  fengsilin 7af60d3214 feat(pricing): 分组价格表格添加缓存读取和缓存创建价格列 1 month ago
  fengsilin 3a78e9c93a fix(user-migration): revalidate drift and verify synced quota 1 month ago
  fengsilin 36b47379f7 test(user-migration): cover rescan and batch status regressions 1 month ago
  fengsilin a406c94b93 feat(user-migration): add root migration dashboard 1 month ago
  fengsilin 7282142316 feat(user-migration): add root admin migration API 1 month ago
  fengsilin 2ecaa4f818 feat(user-migration): add executor and verifier 1 month ago
  fengsilin 9e72c2f068 feat(user-migration): add scan service with conflict analysis 1 month ago
  fengsilin 102b15ad41 feat(user-migration): add internal migration endpoints and client 1 month ago
  fengsilin 38ba0dd24b feat(user-migration): add imported user and oauth migration helpers 1 month ago
  fengsilin 26cb2273a4 feat(user-migration): add batch item and quota grant models 1 month ago
  fengsilin 87a0f94975 refactor: remove user-channel-ratio feature 1 month ago
  fengsilin 0c279e6662 docs: add overseas user migration design 1 month ago
  fengsilin 6ad0524172 test: add comprehensive pricing and channel selection tests 1 month ago
  fengsilin 052b41562a fix(pricing): restore group info on model cards 1 month ago
  fengsilin 0c2ea777f2 fix: remove channel-pricing API calls from frontend 1 month ago
  fengsilin c64e9c0b40 chore: remove .agents from git tracking 1 month ago
  fengsilin a88ab5914a chore: remove web/dist from git tracking 1 month ago
  fengsilin 7a2f584317 fix(claude): use gjson/sjson for reliable thinking.type replacement 1 month ago
  fengsilin 8150993caf fix(admin): improve channel form inputs 1 month ago
  fengsilin 3aa4cdf2ec feat(models): add structured endpoint editor 1 month ago
  fengsilin 481b43cfc5 fix(claude): send adaptive thinking type 1 month ago
  fengsilin 23f9ed4ccd fix(redemption): block synced users from redeeming 1 month ago
  fengsilin 2af3a9301d fix(pricing): remove stale channel pricing backend 1 month ago
  fengsilin 3a6dec8a6c Merge branch 'feat/remove-channel-pricing' 1 month ago
  fengsilin 6479883ddc fix: handle object-type arguments in Responses API stream 1 month ago
  fengsilin 674b96298c refactor(model): remove channel pricing 1 month ago
  fengsilin 415174a676 test(model): switch pricing tests to global defaults 1 month ago
  fengsilin 7ecedd8e57 fix: restore clean go test baseline 1 month ago
  fengsilin 3cf84775f2 feat: add channel metrics collector, fix search pagination and vendor sort order 1 month ago
  fengsilin 6e81faa698 perf: lazy-read body in DoApiRequest for passthrough model mapping 1 month ago
  fengsilin 6434a523a1 feat: passthrough model mapping 1 month ago
  fengsilin 4584a94620 Merge branch 'feat/custom-nav-link' 1 month ago
  fengsilin d32050d081 fix: restore login widget bundle and guard custom nav link 1 month ago
  fengsilin 54573991a4 feat: add customLink configuration UI in header nav settings 1 month ago
  fengsilin 49c6755622 feat: add customLink to navigation hook with filtering logic 1 month ago
  fengsilin 79f8a18cef feat(widget): 重构登录组件为无 UI 的 headless SDK 1 month ago
  fengsilin 60a1403ae2 fix(widget): Dockerfile 添加 Widget 构建步骤 1 month ago
  fengsilin aa15965467 fix: Dockerfile 使用华为云镜像源解决构建网络问题 1 month ago
  fengsilin 1c5e6d03d9 fix: 禁止修改同步用户余额 1 month ago
  fengsilin fe62509b03 feat(metrics): 重试循环和错误处理中采集指标 1 month ago
  fengsilin e51a66c5f9 feat(metrics): postConsumeQuota 中采集 token/配额/延迟指标 1 month ago
  fengsilin 123f0d95ae feat(metrics): 注册 /metrics 端点 1 month ago
  fengsilin 87628f5469 feat(metrics): 渠道状态自定义 Collector 1 month ago
  fengsilin 67e78f2c44 feat(metrics): Gin 中间件 - 活跃请求、延迟、状态码 1 month ago
  fengsilin 9d4378dbb5 feat(metrics): 错误分类函数及测试 1 month ago
  fengsilin 688c5f8666 feat(metrics): Prometheus 指标定义和注册 1 month ago
  fengsilin 5afa4c48fc feat: 用户表新增注册时间字段 + 用户列表显示邮箱和注册时间 1 month ago
  fengsilin 0c369eba51 docs: 添加登录 Widget 使用指南 1 month ago
  fengsilin 763cf86a7a fix(widget): 修复跨域 CORS 问题 1 month ago
  fengsilin de3b359dea feat(widget): 登录 Widget v1 — 密码登录,Shadow DOM 隔离 1 month ago
  fengsilin 3c5b821584 docs: 添加登录 Widget 实施计划 1 month ago
  fengsilin 08b89e6252 docs: 添加登录 Widget 设计文档 1 month ago
  fengsilin 84eb6d6a88 feat: 端点格式限制功能 + 完善测试 1 month ago
  fengsilin 4c76265f32 test: 添加渠道选择删除 E2E 测试 + 清理临时截图 + 新增辅助脚本 1 month ago
  fengsilin c64afdc6c2 fix: 清理渠道选择删除后的代码质量问题 1 month ago
  fengsilin 631e030411 test: 添加渠道选择功能删除后的验证测试 1 month ago
  fengsilin b49842d543 refactor: 清理 Distribute 中的冗余条件判断和缩进 1 month ago
  fengsilin ea0d4e5cde refactor: 完整删除前端渠道选择和默认通道 UI 1 month ago
  fengsilin baed9754d0 refactor: 完整删除渠道选择功能后端代码 1 month ago
  fengsilin 547ed7e307 docs: 渠道选择功能完整删除设计方案 1 month ago
  fengsilin f1a68027f1 Revert "feat(cache): OpenAI→Claude 转换自动注入 prompt caching" 1 month ago
  fengsilin ff14caed7c feat(cache): OpenAI→Claude 转换自动注入 prompt caching 1 month ago
  fengsilin 155aa12014 perf(docker): BuildKit 缓存加速构建 + fix: 模型列表改用 enabled 接口 1 month ago
  fengsilin f07f2c25f4 feat(rate-limit): 用户模型限流配置改用下拉选择模型 1 month ago
  fengsilin 9ee5f535c8 Merge branch 'feat/relay-capture' 1 month ago
  fengsilin 714d8239d1 feat: Relay 请求抓包日志 + 用户模型 RPM 限流 1 month ago
  fengsilin 69ed7d85d7 feat(sidebar): 添加监控外部链接菜单项 2 months ago
  fengsilin bae5c1faa5 feat(pricing): 隐藏模型广场标签和端点类型筛选器 2 months ago
  fengsilin e6b4a06557 feat(logs): 隐藏使用日志中的分组信息展示 2 months ago
  fengsilin f9eeb7e47e fix(billing): 修复 Claude 计费明细未应用用户渠道折扣的问题 2 months ago
  fengsilin f7ff684d0d feat(pricing): 渠道定价展示用户折扣价格 2 months ago
  fengsilin e1b6aff7e1 fix(log): 对普通用户隐藏 Chat ID 和 Upstream ID 2 months ago
  fengsilin 02e0940555 fix(playground): 清空互斥模型默认列表 2 months ago
  fengsilin 2d1a8d2f28 fix: 设置 chat_id + 优化日志查询参数构建 2 months ago
  fengsilin 3e86f523db feat(playground): temperature/top_p 参数互斥设置 2 months ago
  fengsilin d26ad9e635 feat: 渠道-模型级联选择 + 隐藏首页统计卡片 2 months ago
  fengsilin bb708510fe docs: 用户倍率渠道-模型级联选择器设计文档 2 months ago
  fengsilin b026f6bc04 feat: 用户倍率前端展示 + 渠道下拉选择 + 上游调试日志 2 months ago
  fengsilin 9dbf3365fe feat(pricing): 添加用户-模型-渠道倍率功能 2 months ago
  fengsilin 17a0c2c26a refactor(playground): 隐藏对话页面的分组选择框 2 months ago
  fengsilin b033edec56 fix: pending sync records 队列被 quota=0 旧记录阻塞 2 months ago
  fengsilin 8945f8c39e fix: synced 用户剩余额度取 synced_quota + 记录 chat ID 2 months ago
  fengsilin e41f2951f7 fix: preserve relay body and capture invalid responses body 2 months ago
  fengsilin b2cb7bcf60 feat: record chat and upstream ids in usage logs 2 months ago
  fengsilin caf457ce9e feat: 错误时记录上游响应体,支持流式 2 months ago
  fengsilin 575f84064e feat: codex API key 模式 + 前端凭证回显修复 2 months ago
  fengsilin ca91922904 docs: add log chat/upstream id design 2 months ago
  fengsilin 332d5a62f2 feat: support codex api key credentials 2 months ago
  fengsilin 1c0898257c feat: add redemption remarks 2 months ago
  fengsilin f18e51e3a9 feat: add redemption remark support 2 months ago
  fengsilin fae24fdf37 test: stabilize channel affinity usage cache tests 2 months ago
  fengsilin fc0257f6ca chore: ignore local worktrees 2 months ago
  fengsilin d8272e7707 feat(channel): 添加渠道"对外名称"(public_name)字段 2 months ago
  fengsilin 73d10b5799 feat: 错误日志记录上游 request-id 和响应体,Playground 渠道路由改为 header 传递 2 months ago
  fengsilin 5b8303ae81 fix: 登录后跳转来源页,充值金额校验优化 2 months ago
  fengsilin 8bb2883d7e feat(home): 首页定价卡片添加"去体验"按钮,跳转 Playground 2 months ago
  fengsilin a0ac37b2a9 fix: 替换 println 为 SysLog,修复 logger nil context 崩溃 2 months ago
  fengsilin 74cc8c0d56 feat(playground): 添加渠道选择功能,支持指定渠道体验模型 2 months ago
  fengsilin a34d44a824 fix(user): 修复从节点同步用户余额显示为 0 的问题 2 months ago
  fengsilin 967309fb56 feat: 添加邮箱后缀注册额度规则功能 2 months ago
  fengsilin 83f767b40b Merge branch 'worktree-feat-captcha' 2 months ago
  fengsilin d24d68a916 fix(i18n): 修复令牌和充值页面货币显示不一致问题 2 months ago
  fengsilin fe88bc7758 Merge branch 'worktree-feat-captcha' 2 months ago
  fengsilin 1912dd3b72 feat: 添加图片验证码功能,防止脚本批量注册 2 months ago
  fengsilin 20163dbad5 fix(legal): 法律文档页面切换语言时强制重新加载内容 2 months ago
  fengsilin 56e8d14485 fix: 语言切换后 TextArea 内容丢失 2 months ago
  fengsilin 65bd1dabb8 fix: 法律文档编辑器切换语言时 TextArea 内容未刷新 2 months ago
  fengsilin 934b4b10a7 feat: 法律文档双语配置(中/英)+ 精简语言支持至中英双语 2 months ago
  fengsilin 676902ed58 Merge branch 'worktree-feat-terms-usage-policy' 2 months ago
  fengsilin 2757f77011 feat(i18n): 补全服务条款和使用政策的多语言翻译 2 months ago
  fengsilin 65c6446a30 Merge branch 'worktree-feat-terms-usage-policy' 2 months ago
  fengsilin 1cc75aa9a2 feat: 服务条款和使用政策页面,支持后台 Markdown 配置 2 months ago
  fengsilin ed195cada5 fix(frontend): sort_order 默认值显示为"未设置",编辑弹窗增加提示 2 months ago
  fengsilin d5157a779c fix(pricing): 统一 sort_order 排序逻辑,有 meta 但未设置的不优先 2 months ago
  fengsilin 54e53108ec fix(sort): sort_order 默认值改为 999999,简化排序逻辑 2 months ago
  fengsilin 3b99c83e32 fix(sort): sort_order=0 的记录排到最后,非0值按升序排列 2 months ago
  fengsilin 518f1ab87f refactor(frontend): 移除拖拽排序,改为 sort_order 数值输入 2 months ago
  fengsilin 65880fc69b feat(frontend): Vendor Tab 和 Model 表格支持拖拽排序,移除定价页字母排序 2 months ago
  fengsilin 5e0f38d26c chore(frontend): 安装 @dnd-kit 拖拽排序库 2 months ago
  fengsilin e507897f21 feat(pricing): updatePricing 按 sort_order 排序 vendors 和 models 2 months ago
  fengsilin 749064932c feat: 新增 PUT /api/vendors/reorder 和 /api/models/reorder 批量排序接口 2 months ago
  fengsilin f8859b7c35 feat: Vendor 和 Model 添加 sort_order 字段,查询按 sort_order ASC 排序 2 months ago
  fengsilin 3a29f3772f feat(channel): 添加模型默认通道功能,支持优先路由和卡片标识 2 months ago
  fengsilin 93e331b624 feat: 支持 LOGO_FILE_PATH 环境变量指定本地 Logo,优化定价与语言设置 2 months ago
  fengsilin 94f24b29f7 style: 首页文案"全球"改为"顶级" 2 months ago
  fengsilin f0f0215b20 fix(pricing): 修复缓存倍率写入默认值1和缓存创建token丢失 2 months ago
  fengsilin 387f6c1ae0 feat(pricing): 定价数据源切换到渠道表,新增缓存价格展示 2 months ago
  fengsilin 00058cd671 refactor(channel-pricing): 提取辅助方法消除重复代码,修复缓存一致性 2 months ago
  fengsilin 42082311d3 fix: 修复默认语言不生效的问题 2 months ago
  fengsilin 9eb58c684e feat: 后台设置默认语言 2 months ago
  fengsilin a2b27ea2a6 docs: 后台设置默认语言实施计划 2 months ago
  fengsilin 436fcb405b docs: 后台设置默认语言功能设计文档 2 months ago
  fengsilin 8d22c4e80e style(home): 临时隐藏首页工具链/核心价值/工作流/生态伙伴 section 2 months ago
  fengsilin 355757404b merge: feat/channel-pricing-extended → master 2 months ago
  fengsilin 235f7c6e5f refactor(channel-pricing): 提取 ParseTagIds 辅助函数 + 移除调试 console.log 2 months ago
  fengsilin f3577590bc test(channel-pricing): 添加 API 端到端集成测试脚本 2 months ago
  fengsilin 129e17ac42 feat(channel-pricing): 前端支持缓存/图片/音频倍率编辑和展示 2 months ago
  fengsilin c84c26b5a2 feat(channel-pricing): GetChannelPricingByModelWithChannelInfo 支持 CASE WHEN 回退 + 扩展字段 2 months ago
  fengsilin d5d4714908 test(channel-pricing): 缓存写穿 + 字段默认值单元测试 2 months ago
  fengsilin 646f37dd4b feat(channel-pricing): Controller 扩展 API 支持新字段 + 输入校验 + 操作日志 2 months ago
  fengsilin 3a24fd598f feat(channel-pricing): ModelPriceHelper + UpdatePriceDataForChannelPricing 适配新签名,支持扩展比率覆盖 2 months ago
  fengsilin e5f029f91e feat(channel-pricing): InitDB 启动时全量加载渠道定价缓存 2 months ago
  fengsilin 2d4c73d3aa refactor(channel-pricing): 结构体新增扩展字段 + 缓存重写为全量加载写穿 2 months ago
  fengsilin 156618fdad merge: feat/alipay-payment → master 2 months ago
  fengsilin b514a0798b refactor(payment): 合并微信/支付宝支付重复代码,删除冗余文件 2 months ago
100 changed files with 8453 additions and 785 deletions
Split View
  1. +21
    -1
      .dockerignore
  2. +15
    -0
      .env.dev
  3. +1
    -0
      .gitattributes
  4. +15
    -0
      .gitignore
  5. +7
    -0
      .plans/alipay-payment/backend-dev/findings.md
  6. +7
    -0
      .plans/alipay-payment/backend-dev/progress.md
  7. +7
    -0
      .plans/alipay-payment/backend-dev/task-alipay-backend/findings.md
  8. +7
    -0
      .plans/alipay-payment/backend-dev/task-alipay-backend/progress.md
  9. +46
    -0
      .plans/alipay-payment/backend-dev/task-alipay-backend/task_plan.md
  10. +23
    -0
      .plans/alipay-payment/backend-dev/task_plan.md
  11. +26
    -0
      .plans/alipay-payment/decisions.md
  12. +7
    -0
      .plans/alipay-payment/findings.md
  13. +7
    -0
      .plans/alipay-payment/frontend-dev/findings.md
  14. +7
    -0
      .plans/alipay-payment/frontend-dev/progress.md
  15. +7
    -0
      .plans/alipay-payment/frontend-dev/task-alipay-frontend/findings.md
  16. +7
    -0
      .plans/alipay-payment/frontend-dev/task-alipay-frontend/progress.md
  17. +36
    -0
      .plans/alipay-payment/frontend-dev/task-alipay-frontend/task_plan.md
  18. +21
    -0
      .plans/alipay-payment/frontend-dev/task_plan.md
  19. +16
    -0
      .plans/alipay-payment/progress.md
  20. +7
    -0
      .plans/alipay-payment/reviewer/findings.md
  21. +5
    -0
      .plans/alipay-payment/reviewer/progress.md
  22. +16
    -0
      .plans/alipay-payment/reviewer/task_plan.md
  23. +47
    -0
      .plans/alipay-payment/task_plan.md
  24. +569
    -0
      .plans/channel-public-name.md
  25. +72
    -0
      .superpowers/brainstorm/6600-1777284109/content/button-position.html
  26. +3
    -0
      .superpowers/brainstorm/6600-1777284109/content/waiting.html
  27. +1
    -0
      .superpowers/brainstorm/6600-1777284109/state/server-stopped
  28. +1
    -0
      .superpowers/brainstorm/6600-1777284109/state/server.pid
  29. +6
    -5
      Dockerfile
  30. +123
    -0
      alan-claudecode-429-upstream-ids.txt
  31. BIN
      atlas-edit-test.png
  32. +23
    -0
      atlas-gemini-image.json
  33. BIN
      atlascloud-local-test.png
  34. +4
    -0
      common/api_type.go
  35. +30
    -0
      common/captcha.go
  36. +58
    -0
      common/codex_credential.go
  37. +10
    -0
      common/constants.go
  38. +1
    -0
      common/endpoint_defaults.go
  39. +2
    -0
      common/endpoint_type.go
  40. +24
    -4
      common/limiter/limiter.go
  41. +18
    -0
      common/limiter/lua/rate_refund.lua
  42. +64
    -0
      common/metrics/collector.go
  43. +93
    -0
      common/metrics/collector_test.go
  44. +53
    -0
      common/metrics/errorclassify.go
  45. +85
    -0
      common/metrics/errorclassify_test.go
  46. +119
    -0
      common/metrics/metrics.go
  47. +242
    -0
      common/metrics/metrics_test.go
  48. +28
    -0
      common/metrics/middleware.go
  49. +12
    -0
      common/rate-limit.go
  50. +18
    -0
      common/str.go
  51. +21
    -0
      common/tianyiyun_channel_test.go
  52. +4
    -0
      common/utils.go
  53. +27
    -0
      common/utils_test.go
  54. +118
    -109
      constant/channel.go
  55. +2
    -1
      constant/context_key.go
  56. +21
    -8
      constant/endpoint_type.go
  57. +2
    -0
      controller/channel-test.go
  58. +23
    -38
      controller/channel.go
  59. +0
    -325
      controller/channel_pricing.go
  60. +36
    -0
      controller/channel_test_tianyiyun_test.go
  61. +111
    -0
      controller/codex_channel_test.go
  62. +8
    -10
      controller/codex_usage.go
  63. +158
    -0
      controller/doubao_aiping_video.go
  64. +209
    -0
      controller/doubao_aiping_video_test.go
  65. +203
    -0
      controller/doubao_asset.go
  66. +103
    -0
      controller/doubao_asset_test.go
  67. +128
    -0
      controller/email_quota_rule.go
  68. +264
    -0
      controller/email_quota_rule_test.go
  69. +203
    -0
      controller/internal_user_migration.go
  70. +491
    -0
      controller/internal_user_migration_test.go
  71. +510
    -0
      controller/kling_aiping_native.go
  72. +250
    -0
      controller/kling_aiping_native_test.go
  73. +10
    -2
      controller/log.go
  74. +94
    -0
      controller/log_identity_test.go
  75. +74
    -5
      controller/misc.go
  76. +59
    -0
      controller/model_display_pricing_controller.go
  77. +96
    -0
      controller/model_display_pricing_controller_test.go
  78. +8
    -0
      controller/model_meta.go
  79. +89
    -0
      controller/model_meta_test.go
  80. +59
    -0
      controller/model_pricing_controller.go
  81. +153
    -0
      controller/model_pricing_controller_test.go
  82. +10
    -0
      controller/playground.go
  83. +171
    -0
      controller/pricing.go
  84. +0
    -131
      controller/pricing_tag.go
  85. +349
    -0
      controller/pricing_user_test.go
  86. +2
    -0
      controller/redemption.go
  87. +159
    -0
      controller/redemption_test.go
  88. +229
    -38
      controller/relay.go
  89. +152
    -0
      controller/relay_test.go
  90. +55
    -0
      controller/reorder.go
  91. +42
    -20
      controller/topup.go
  92. +10
    -63
      controller/topup_alipay.go
  93. +2
    -24
      controller/topup_wechat.go
  94. +28
    -1
      controller/user.go
  95. +376
    -0
      controller/user_migration.go
  96. +839
    -0
      controller/user_migration_test.go
  97. +74
    -0
      controller/user_rate_limit.go
  98. +152
    -0
      controller/user_rate_limit_test.go
  99. +146
    -0
      controller/user_sync_guard_test.go
  100. +106
    -0
      controller/user_topup_test.go

+ 21
- 1
.dockerignore View File

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

+ 15
- 0
.env.dev View File

@@ -0,0 +1,15 @@
# 本地开发环境变量,与 docker/dev/docker-compose.yml 配套
# 用法(WSL / bash):
# source .env.dev && go run main.go
# 用法(PowerShell):
# $env:SQL_DSN="root:Lanqi123456@tcp(localhost:33306)/new-api"; $env:REDIS_CONN_STRING="redis://localhost:36379"; go run main.go

export SQL_DSN="root:Lanqi123456@tcp(localhost:33306)/new-api"
export REDIS_CONN_STRING="redis://localhost:36379"
export SESSION_SECRET="dev-secret-do-not-use-in-prod"
export SESSION_NAME="session_dev"
export DEBUG=1
export ERROR_LOG_ENABLED=true
export BATCH_UPDATE_ENABLED=true
export TZ=Asia/Shanghai
export ATLASCLOUD_API_KEY=apikey-ba28bd46b2bf470aa8a6718db4b6839a

+ 1
- 0
.gitattributes View File

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

+ 15
- 0
.gitignore View File

@@ -23,9 +23,24 @@ plans
docs/plans/
CLAUDE.md
.claude
.agents/
.worktrees/
logs/
docs/superpowers

# Runtime request/response capture files
/[0-9]*.json

# Local dev server logs
web/vite-*.log

# E2E / Playwright artifacts
test-artifacts/
web/.playwright/
web/playwright-report/
web/test-results/
web/e2e/.auth/

electron/node_modules
electron/dist
data/


+ 7
- 0
.plans/alipay-payment/backend-dev/findings.md View File

@@ -0,0 +1,7 @@
# backend-dev - 发现索引

> 纯索引——每个条目应简短(Status + Report 链接 + Summary)。

---

<初始为空,工作中填写>

+ 7
- 0
.plans/alipay-payment/backend-dev/progress.md View File

@@ -0,0 +1,7 @@
# backend-dev - 工作日志

> 用于上下文恢复。压缩/重启后先读此文件。

---

<初始为空,工作中填写>

+ 7
- 0
.plans/alipay-payment/backend-dev/task-alipay-backend/findings.md View File

@@ -0,0 +1,7 @@
# 支付宝后端 - 发现记录

> 此任务开发中的技术发现。

---

<初始为空>

+ 7
- 0
.plans/alipay-payment/backend-dev/task-alipay-backend/progress.md View File

@@ -0,0 +1,7 @@
# 支付宝后端 - 工作日志

> 上下文恢复时只需读此文件。

---

<初始为空>

+ 46
- 0
.plans/alipay-payment/backend-dev/task-alipay-backend/task_plan.md View File

@@ -0,0 +1,46 @@
# 支付宝后端 - 任务计划

> 所属智能体: backend-dev
> 状态: pending
> 创建: 2026-04-14

## 目标

实现支付宝当面付(扫码支付)的完整后端功能,包括配置层、控制器、模型和路由。

## 详细步骤

- [ ] 1. 新建 `setting/payment_alipay.go`:配置变量 + IsAlipayConfigured() + OnAlipayConfigChanged
- [ ] 2. 修改 `model/option.go`:InitOptionMap + updateOptionMap + triggerAlipayReset
- [ ] 3. 新建 `model/topup_alipay.go`:RechargeAlipay 函数(事务+行锁+幂等)
- [ ] 4. 新建 `controller/topup_alipay.go`:客户端缓存 + 5 个 HTTP 处理函数
- [ ] 5. 修改 `controller/topup.go`:GetTopUpInfo 添加支付宝支付方式
- [ ] 6. 修改 `model/topup.go`:ManualCompleteTopUp 添加 "alipay"
- [ ] 7. 修改 `router/api-router.go`:注册 4 个路由
- [ ] 8. `go build` 编译验证

## 涉及文件

- `setting/payment_alipay.go` — 新建:支付宝配置
- `setting/payment_wechat.go` — 参考:微信支付配置模式
- `controller/topup_alipay.go` — 新建:支付宝控制器
- `controller/topup_wechat.go` — 参考:微信支付控制器
- `model/topup_alipay.go` — 新建:支付宝充值模型
- `model/topup_wechat.go` — 参考:微信充值模型
- `model/option.go` — 修改:注册配置 key
- `model/topup.go` — 修改:添加 alipay 支付方式
- `controller/topup.go` — 修改:GetTopUpInfo
- `router/api-router.go` — 修改:注册路由

## 关键技术点

- gopay alipay 子包:`github.com/go-pay/gopay/alipay`
- TradePrecreate API:当面付预下单
- 金额单位:元(字符串 "7.00"),不是微信的分
- 回调验签:alipay.VerifySign(alipayPublicKey, notifyReq)
- 回调响应:纯文本 "success"
- 订单号前缀:ali(区分微信的 wx)

## 依赖

- 无外部依赖,go-pay 已在项目中

+ 23
- 0
.plans/alipay-payment/backend-dev/task_plan.md View File

@@ -0,0 +1,23 @@
# backend-dev - 任务计划

> 角色: 后端开发
> 状态: pending
> 分配的任务: 支付宝当面付后端完整实现

## 任务

- [ ] 步骤 1: 新建 setting/payment_alipay.go — 支付宝配置变量
- [ ] 步骤 2: 修改 model/option.go — 注册配置 key + 热更新
- [ ] 步骤 3: 新建 model/topup_alipay.go — 充值完成处理逻辑
- [ ] 步骤 4: 新建 controller/topup_alipay.go — 核心控制器
- [ ] 步骤 5: 修改 controller/topup.go — GetTopUpInfo 添加支付宝
- [ ] 步骤 6: 修改 model/topup.go — ManualCompleteTopUp 添加 alipay
- [ ] 步骤 7: 修改 router/api-router.go — 注册路由
- [ ] 步骤 8: 编译验证

## 备注

- 参考文件:controller/topup_wechat.go(微信支付控制器)、setting/payment_wechat.go(配置)、model/topup_wechat.go(模型)
- 使用 gopay v1.5.117 的 alipay 子包
- 关键差异:金额单位是元(不是分)、回调响纯文本 "success"(不是 JSON)
- 完成后找 reviewer 审查

+ 26
- 0
.plans/alipay-payment/decisions.md View File

@@ -0,0 +1,26 @@
# alipay-payment - 架构决策记录

> 记录每个决策及其理由。

---

## D1: 支付产品选择

- 日期: 2026-04-14
- 决策: 使用支付宝当面付(扫码支付)
- 理由: 与现有微信 Native 支付模式一致,可最大程度复用架构
- 考虑过的替代方案: PC 网站支付、H5 手机支付

## D2: 签名方式

- 日期: 2026-04-14
- 决策: 公钥模式(RSA2)
- 理由: 参数少,配置简单,只需 AppID + 应用私钥 + 支付宝公钥
- 考虑过的替代方案: 证书模式(更安全但配置复杂)

## D3: 团队配置

- 日期: 2026-04-14
- 决策: 3 角色(backend-dev + frontend-dev + reviewer)
- 理由: 前后端可并行开发,支付涉及资金安全需要代码审查
- 考虑过的替代方案: 2 角色(无审查)、1 角色(全栈)

+ 7
- 0
.plans/alipay-payment/findings.md View File

@@ -0,0 +1,7 @@
# alipay-payment - 发现与技术记录

> 由团队智能体自动更新。每条标注来源。

---

<工作中添加条目>

+ 7
- 0
.plans/alipay-payment/frontend-dev/findings.md View File

@@ -0,0 +1,7 @@
# frontend-dev - 发现索引

> 纯索引——每个条目应简短。

---

<初始为空,工作中填写>

+ 7
- 0
.plans/alipay-payment/frontend-dev/progress.md View File

@@ -0,0 +1,7 @@
# frontend-dev - 工作日志

> 用于上下文恢复。

---

<初始为空,工作中填写>

+ 7
- 0
.plans/alipay-payment/frontend-dev/task-alipay-frontend/findings.md View File

@@ -0,0 +1,7 @@
# 支付宝前端 - 发现记录

> 此任务开发中的技术发现。

---

<初始为空>

+ 7
- 0
.plans/alipay-payment/frontend-dev/task-alipay-frontend/progress.md View File

@@ -0,0 +1,7 @@
# 支付宝前端 - 工作日志

> 上下文恢复时只需读此文件。

---

<初始为空>

+ 36
- 0
.plans/alipay-payment/frontend-dev/task-alipay-frontend/task_plan.md View File

@@ -0,0 +1,36 @@
# 支付宝前端 - 任务计划

> 所属智能体: frontend-dev
> 状态: pending
> 创建: 2026-04-14

## 目标

实现支付宝当面付的前端功能,包括管理后台设置页、用户充值页改造、二维码模态框重构。

## 详细步骤

- [ ] 1. 重构 `WechatPayQRCodeModal.jsx` → 通用 `QRCodePayModal`(新增 title/subtitle/statusApiPath props)
- [ ] 2. 新建 `SettingsPaymentGatewayAlipay.jsx`:管理后台支付宝设置页
- [ ] 3. 修改 `PaymentSetting.jsx`:导入支付宝设置组件 + getOptions 解析 + 渲染
- [ ] 4. 修改 `topup/index.jsx`:支付宝状态变量 + API 调用 + 模态框渲染
- [ ] 5. 修改 `RechargeCard.jsx`:传递 enableAlipayTopUp prop
- [ ] 6. 更新 i18n 翻译文件
- [ ] 7. `bun run build` 验证编译通过

## 涉及文件

- `web/src/components/topup/WechatPayQRCodeModal.jsx` — 重构为通用模态框
- `web/src/pages/Setting/Payment/SettingsPaymentGatewayAlipay.jsx` — 新建
- `web/src/components/settings/PaymentSetting.jsx` — 修改
- `web/src/components/topup/index.jsx` — 修改
- `web/src/components/topup/RechargeCard.jsx` — 修改(最小)
- `web/src/i18n/locales/en.json` — 修改(翻译)

## 关键技术点

- QRCodePayModal 需要兼容微信和支付宝两种场景
- 支付宝轮询 API: /api/user/alipay/pay/status
- 支付宝创建订单 API: POST /api/user/alipay/pay
- 支付宝图标 SiAlipay 已存在于 RechargeCard.jsx
- i18n key 遵循现有命名规范

+ 21
- 0
.plans/alipay-payment/frontend-dev/task_plan.md View File

@@ -0,0 +1,21 @@
# frontend-dev - 任务计划

> 角色: 前端开发
> 状态: pending
> 分配的任务: 支付宝当面付前端完整实现

## 任务

- [ ] 步骤 1: 重构 WechatPayQRCodeModal.jsx 为通用 QRCodePayModal
- [ ] 步骤 2: 新建 SettingsPaymentGatewayAlipay.jsx 管理后台设置页
- [ ] 步骤 3: 修改 PaymentSetting.jsx 注册支付宝设置
- [ ] 步骤 4: 修改 topup/index.jsx 添加支付宝支付逻辑
- [ ] 步骤 5: 修改 RechargeCard.jsx 传递 enableAlipayTopUp
- [ ] 步骤 6: 更新 i18n 翻译文件
- [ ] 步骤 7: bun run build 验证

## 备注

- 参考文件:WechatPayQRCodeModal.jsx(二维码模态框)、SettingsPaymentGatewayWechat.jsx(设置页)
- 支付宝图标 SiAlipay 已存在于 RechargeCard.jsx
- 完成后找 reviewer 审查

+ 16
- 0
.plans/alipay-payment/progress.md View File

@@ -0,0 +1,16 @@
# alipay-payment - 进度日志

> 按时间线记录。每条记录谁做了什么。

---

## 2026-04-14 Session 1 — 团队搭建

### 已完成
- [x] 读取设计文档
- [x] 确认团队配置:3 角色(backend-dev, frontend-dev, reviewer)
- [x] 创建规划文件和目录结构

### 待办
- [ ] 启动团队成员
- [ ] 开始并行开发

+ 7
- 0
.plans/alipay-payment/reviewer/findings.md View File

@@ -0,0 +1,7 @@
# reviewer - 发现索引

> 纯索引。

---

<初始为空,工作中填写>

+ 5
- 0
.plans/alipay-payment/reviewer/progress.md View File

@@ -0,0 +1,5 @@
# reviewer - 工作日志

---

<初始为空>

+ 16
- 0
.plans/alipay-payment/reviewer/task_plan.md View File

@@ -0,0 +1,16 @@
# reviewer - 任务计划

> 角色: 代码审查
> 状态: pending
> 分配的任务: 等待 backend-dev 和 frontend-dev 完成后进行代码审查

## 任务

- [ ] 审查 backend-dev 支付宝后端代码(安全 + 质量)
- [ ] 审查 frontend-dev 支付宝前端代码(质量 + 体验)

## 备注

- 支付涉及资金安全,重点关注:签名验签、金额处理、幂等性、并发安全
- 后端审查重点:回调验签、金额单位(元 vs 分)、订单幂等
- 前端审查重点:二维码模态框重构兼容性、支付状态轮询

+ 47
- 0
.plans/alipay-payment/task_plan.md View File

@@ -0,0 +1,47 @@
# alipay-payment - 主计划

> 状态: PLANNING
> 创建: 2026-04-14
> 更新: 2026-04-14
> 团队: alipay-payment (backend-dev, frontend-dev, reviewer)
> 决策记录: .plans/alipay-payment/decisions.md

---

## 1. 项目概述

在 New-API 项目中集成支付宝当面付(扫码支付),复用微信支付架构模式。使用 go-pay 框架,公钥模式签名,仅支持充值业务。

设计文档:`D:\markdown\workspace-lanqi\new-api\2026-04-14\支付宝当面付集成设计.md`

---

## 2. 文档索引

| 文档 | 位置 | 内容 |
|------|------|------|
| 设计文档(Obsidian) | D:\markdown\workspace-lanqi\new-api\2026-04-14\支付宝当面付集成设计.md | 完整设计方案 |

---

## 3. 阶段概览

- 阶段 1: 后端开发 — backend-dev 实现配置层、控制器、模型、路由
- 阶段 2: 前端开发 — frontend-dev 实现管理设置页、充值页改造、二维码模态框重构(可与阶段 1 并行)
- 阶段 3: 代码审查 — reviewer 审查后端和前端代码

---

## 4. 任务汇总

| # | 任务 | 负责人 | 状态 | 计划文件 |
|---|------|--------|------|----------|
| T1 | 后端支付宝支付完整实现 | backend-dev | pending | .plans/alipay-payment/backend-dev/task-alipay-backend/ |
| T2 | 前端支付宝支付完整实现 | frontend-dev | pending | .plans/alipay-payment/frontend-dev/task-alipay-frontend/ |
| T3 | 代码审查 | reviewer | pending | .plans/alipay-payment/reviewer/ |

---

## 5. 当前阶段

准备启动阶段 1 和阶段 2(并行开发)。

+ 569
- 0
.plans/channel-public-name.md View File

@@ -0,0 +1,569 @@
# Channel PublicName 实施计划

> **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:** 为 Channel 添加 `public_name` 字段,让管理员可以为每个渠道设置用户可见的友好名称,替代前端硬编码的"通道一/二/三"。

**Architecture:** 后端 Channel 模型新增字段 → AutoMigrate 自动建列 → 迁移函数回填数据 → 定价查询和 API 返回公共名称 → 前端表单支持编辑 → 展示层使用公共名称。

**Tech Stack:** Go (GORM) / React (Semi Design) / SQLite+MySQL+PostgreSQL

---

### Task 1: Channel 模型添加 PublicName 字段

**Files:**
- Modify: `model/channel.go:28`

- [ ] **Step 1: 在 Channel 结构体 Name 字段后添加 PublicName**

当前代码(`model/channel.go:28`):
```go
Name string `json:"name" gorm:"index"`
Weight *uint `json:"weight" gorm:"default:0"`
```

改为:
```go
Name string `json:"name" gorm:"index"`
PublicName string `json:"public_name" gorm:"size:255;default:''"`
Weight *uint `json:"weight" gorm:"default:0"`
```

- [ ] **Step 2: 验证编译通过**

Run: `cd D:/code/new-api && go build ./model/...`
Expected: 编译成功,无错误

- [ ] **Step 3: Commit**

```bash
git add model/channel.go
git commit -m "feat: add PublicName field to Channel model"
```

---

### Task 2: 修改 GetAllChannelsForBinding 查询包含 public_name

**Files:**
- Modify: `model/channel.go:282`

- [ ] **Step 1: 修改 Select 列**

当前代码(`model/channel.go:282`):
```go
err := DB.Select("id, name, type, remark").
```

改为:
```go
err := DB.Select("id, name, public_name, type, remark").
```

- [ ] **Step 2: 验证编译通过**

Run: `cd D:/code/new-api && go build ./model/...`
Expected: 编译成功

- [ ] **Step 3: Commit**

```bash
git add model/channel.go
git commit -m "feat: include public_name in GetAllChannelsForBinding query"
```

---

### Task 3: ChannelPricingWithChannel 扩展 + SQL 查询修改

**Files:**
- Modify: `model/channel_pricing.go:254` (结构体)
- Modify: `model/channel_pricing.go:299` (SELECT)
- Modify: `model/channel_pricing.go:319` (GROUP BY)

- [ ] **Step 1: 结构体添加 ChannelPublicName 字段**

当前代码(`model/channel_pricing.go:254`):
```go
ChannelName string `json:"channel_name"`
ChannelType int `json:"channel_type"`
```

改为:
```go
ChannelName string `json:"channel_name"`
ChannelPublicName string `json:"channel_public_name"`
ChannelType int `json:"channel_type"`
```

- [ ] **Step 2: SELECT 子句添加 channels.public_name**

当前代码(`model/channel_pricing.go:299`):
```go
Select(`abilities.channel_id, channels.name as channel_name, channels.type as channel_type,
```

改为:
```go
Select(`abilities.channel_id, channels.name as channel_name, channels.public_name as channel_public_name, channels.type as channel_type,
```

- [ ] **Step 3: GROUP BY 子句添加 channels.public_name**

当前代码(`model/channel_pricing.go:319`):
```go
Group("abilities.channel_id, channels.name, channels.type, channel_pricings.quota_type,
```

改为:
```go
Group("abilities.channel_id, channels.name, channels.public_name, channels.type, channel_pricings.quota_type,
```

注意:`channels.public_name` 插入在 `channels.name` 后面。

- [ ] **Step 4: 验证编译通过**

Run: `cd D:/code/new-api && go build ./model/...`
Expected: 编译成功

- [ ] **Step 5: Commit**

```bash
git add model/channel_pricing.go
git commit -m "feat: add public_name to channel pricing query"
```

---

### Task 4: pricing.go 默认通道名称使用 public_name

**Files:**
- Modify: `model/pricing.go:290-303`

- [ ] **Step 1: 扩展匿名结构体**

当前代码(`model/pricing.go:290-293`):
```go
var allCPs []struct {
ChannelPricing
ChannelName string
}
```

改为:
```go
var allCPs []struct {
ChannelPricing
ChannelName string
ChannelPublicName string
}
```

- [ ] **Step 2: 扩展 Select 子句**

当前代码(`model/pricing.go:295`):
```go
Select("channel_pricings.*, channels.name as channel_name").
```

改为:
```go
Select("channel_pricings.*, channels.name as channel_name, channels.public_name as channel_public_name").
```

- [ ] **Step 3: 修改 channelNameMap 构建逻辑**

当前代码(`model/pricing.go:303`):
```go
channelNameMap[allCPs[i].ChannelId] = allCPs[i].ChannelName
```

改为:
```go
name := allCPs[i].ChannelPublicName
if name == "" {
name = allCPs[i].ChannelName
}
channelNameMap[allCPs[i].ChannelId] = name
```

- [ ] **Step 4: 验证编译通过**

Run: `cd D:/code/new-api && go build ./model/...`
Expected: 编译成功

- [ ] **Step 5: Commit**

```bash
git add model/pricing.go
git commit -m "feat: use public_name as default channel display name"
```

---

### Task 5: 数据回填迁移函数

**Files:**
- Modify: `model/main.go:306-307` (调用位置)
- Modify: `model/main.go:695` (函数定义位置)

- [ ] **Step 1: 在 migrateDB 的 return nil 之前插入调用**

当前代码(`model/main.go:304-308`):
```go
// 将现有 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)
return nil
}
```

改为:
```go
// 将现有 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
}
```

- [ ] **Step 2: 在文件末尾(第 696 行后)添加迁移函数**

```go
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))
}
}
```

**说明:**
- AutoMigrate(`migrateDB()` 第 262 行 `&Channel{}`)会先创建列
- WHERE 条件保证幂等
- `gorm.Expr("name")` 引用列名,SQLite/MySQL/PostgreSQL 通用
- `model/main.go` 已有 `fmt`、`gorm`、`common` 的 import,无需额外导入

- [ ] **Step 3: 验证编译通过**

Run: `cd D:/code/new-api && go build ./model/...`
Expected: 编译成功

- [ ] **Step 4: Commit**

```bash
git add model/main.go
git commit -m "feat: add migration to backfill public_name from name"
```

---

### Task 6: controller 层验证 + API 返回

**Files:**
- Modify: `controller/channel.go:582-585` (验证)
- Modify: `controller/channel.go:2110-2115` (API 返回)

- [ ] **Step 1: 在 validateChannel 的 isAdd 块中添加 public_name 校验**

当前代码(`controller/channel.go:582-585`):
```go
if isAdd {
if channel == nil || channel.Key == "" {
return fmt.Errorf("channel cannot be empty")
}

// 检查模型名称长度是否超过 255
```

改为:
```go
if isAdd {
if channel == nil || channel.Key == "" {
return fmt.Errorf("channel cannot be empty")
}

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

// 检查模型名称长度是否超过 255
```

**说明:** 保持英文错误消息与现有代码一致。

- [ ] **Step 2: 在 GetUserChannelsForBinding 返回值中添加 public_name**

当前代码(`controller/channel.go:2110-2115`):
```go
result = append(result, gin.H{
"id": ch.Id,
"name": ch.Name,
"type": ch.Type,
"remark": ch.Remark,
})
```

改为:
```go
result = append(result, gin.H{
"id": ch.Id,
"name": ch.Name,
"public_name": ch.PublicName,
"type": ch.Type,
"remark": ch.Remark,
})
```

- [ ] **Step 3: 验证编译通过**

Run: `cd D:/code/new-api && go build ./controller/...`
Expected: 编译成功

- [ ] **Step 4: Commit**

```bash
git add controller/channel.go
git commit -m "feat: validate public_name on channel creation and return in binding API"
```

---

### Task 7: 前端 i18n 翻译

**Files:**
- Modify: `web/src/i18n/locales/zh-CN.json`
- Modify: `web/src/i18n/locales/en.json`

- [ ] **Step 1: zh-CN.json 添加翻译**

在 `"请为渠道命名"` 条目(第 2454 行)之后添加:

```json
"对外名称": "对外名称",
"用户看到的渠道名称,如"标准通道"、"高速通道"": "用户看到的渠道名称,如"标准通道"、"高速通道"",
"请填写对外名称": "请填写对外名称",
"请填写渠道名称、对外名称和渠道密钥!": "请填写渠道名称、对外名称和渠道密钥!",
```

注意:zh-CN.json 的 key 和 value 相同(中文 → 中文)。

- [ ] **Step 2: en.json 添加翻译**

在 `"Please name the channel"` 条目(第 2473 行)之后添加:

```json
"对外名称": "Public Name",
"用户看到的渠道名称,如"标准通道"、"高速通道"": "User-facing channel name, e.g. \"Standard\", \"Fast\"",
"请填写对外名称": "Please enter a public name",
"请填写渠道名称、对外名称和渠道密钥!": "Please enter channel name, public name and key!",
```

- [ ] **Step 3: Commit**

```bash
git add web/src/i18n/locales/zh-CN.json web/src/i18n/locales/en.json
git commit -m "feat: add i18n translations for channel public_name"
```

---

### Task 8: EditChannelModal 表单添加对外名称字段

**Files:**
- Modify: `web/src/components/table/channels/modals/EditChannelModal.jsx:142` (originInputs)
- Modify: `web/src/components/table/channels/modals/EditChannelModal.jsx:1980` (表单 UI)
- Modify: `web/src/components/table/channels/modals/EditChannelModal.jsx:1329` (提交校验)

- [ ] **Step 1: originInputs 添加 public_name 默认值**

当前代码(第 142-143 行):
```javascript
const originInputs = {
name: '',
type: 1,
```

改为:
```javascript
const originInputs = {
name: '',
public_name: '',
type: 1,
```

- [ ] **Step 2: 在 name 的 Form.Input 后添加 public_name 输入框**

当前代码(第 1972-1980 行):
```jsx
<Form.Input
field='name'
label={t('名称')}
placeholder={t('请为渠道命名')}
rules={[{ required: true, message: t('请为渠道命名') }]}
showClear
onChange={(value) => handleInputChange('name', value)}
autoComplete='new-password'
/>

{inputs.type === 33 && (
```

在第 1980 行 `/>` 和第 1982 行 `{inputs.type === 33` 之间插入:

```jsx
<Form.Input
field='name'
label={t('名称')}
placeholder={t('请为渠道命名')}
rules={[{ required: true, message: t('请为渠道命名') }]}
showClear
onChange={(value) => handleInputChange('name', value)}
autoComplete='new-password'
/>
<Form.Input
field='public_name'
label={t('对外名称')}
placeholder={t('用户看到的渠道名称,如"标准通道"、"高速通道"')}
rules={!isEdit ? [{ required: true, message: t('请填写对外名称') }] : []}
showClear
onChange={(value) => handleInputChange('public_name', value)}
autoComplete='new-password'
/>

{inputs.type === 33 && (
```

**说明:** `!isEdit` 条件确保新建时必填、编辑时允许清空。

- [ ] **Step 3: 修改提交前校验**

当前代码(第 1329-1332 行):
```javascript
if (!isEdit && (!localInputs.name || !localInputs.key)) {
showInfo(t('请填写渠道名称和渠道密钥!'));
return;
}
```

改为:
```javascript
if (!isEdit && (!localInputs.name || !localInputs.public_name || !localInputs.key)) {
showInfo(t('请填写渠道名称、对外名称和渠道密钥!'));
return;
}
```

- [ ] **Step 4: Commit**

```bash
git add web/src/components/table/channels/modals/EditChannelModal.jsx
git commit -m "feat: add public_name field to channel edit form"
```

---

### Task 9: ChannelPricingCard 展示公共名称

**Files:**
- Modify: `web/src/components/table/model-pricing/modal/components/ChannelPricingCard.jsx:107`

- [ ] **Step 1: 修改 channelName 赋值逻辑**

当前代码(第 104-108 行):
```javascript
const tableData = channelPricingData.map((item, index) => ({
key: item.channel_id || index,
channelId: item.channel_id,
channelName: `通道${['一', '二', '三', '四', '五', '六', '七', '八', '九', '十'][index] || ` ${index + 1}`}`,
channelTags: item.tags || [],
```

改为:
```javascript
const tableData = channelPricingData.map((item, index) => ({
key: item.channel_id || index,
channelId: item.channel_id,
channelName: item.channel_public_name || ('通道' + (['一', '二', '三', '四', '五', '六', '七', '八', '九', '十'][index] || (index + 1))),
channelTags: item.tags || [],
```

**说明:** 优先使用后端 `channel_public_name`,为空时回退到中文数字。同时修复原代码反引号嵌套语法问题。

- [ ] **Step 2: Commit**

```bash
git add web/src/components/table/model-pricing/modal/components/ChannelPricingCard.jsx
git commit -m "feat: display public_name in ChannelPricingCard"
```

---

### Task 10: 端到端验证

**Files:** 无代码修改,纯验证

- [ ] **Step 1: 后端编译**

Run: `cd D:/code/new-api && go build -o new-api main.go`
Expected: 编译成功

- [ ] **Step 2: 启动服务**

Run: `cd D:/code/new-api && go run main.go`
Expected: 日志中出现 `[Migration] migrateChannelPublicName: backfilled N channels`

- [ ] **Step 3: 再次重启确认幂等**

重启服务后 Expected: 日志中不再出现 backfilled(RowsAffected=0)

- [ ] **Step 4: API 测试 — 创建渠道不带 public_name**

Run:
```bash
curl -s -X POST http://localhost:3000/api/channel/ \
-H "Authorization: Bearer $ADMIN_KEY" \
-H "Content-Type: application/json" \
-d '{"mode":"single","channel":{"type":1,"name":"测试","key":"sk-test"}}'
```
Expected: 返回错误 `"public name cannot be empty"`

- [ ] **Step 5: API 测试 — 创建渠道带 public_name**

Run:
```bash
curl -s -X POST http://localhost:3000/api/channel/ \
-H "Authorization: Bearer $ADMIN_KEY" \
-H "Content-Type: application/json" \
-d '{"mode":"single","channel":{"type":1,"name":"测试","public_name":"标准通道","key":"sk-test"}}'
```
Expected: 返回成功

- [ ] **Step 6: API 测试 — 定价接口包含 public_name**

Run:
```bash
curl -s http://localhost:3000/api/channel-pricing/model/gpt-4
```
Expected: 响应中包含 `channel_public_name` 字段

- [ ] **Step 7: 前端验证**

在浏览器中依次验证:
1. 渠道管理 → 新建渠道 → 可见"对外名称"输入框
2. 不填"对外名称"提交 → 显示验证错误
3. 填写后提交成功
4. 编辑渠道 → 可见已保存的对外名称
5. 编辑时清空对外名称并保存 → 成功(向后兼容)
6. 模型详情侧边栏 → ChannelPricingCard 显示对外名称
7. 首页定价卡片 → "默认通道"标签显示对外名称
8. 未设 public_name 的渠道 → 显示"通道一/二/三..."

+ 72
- 0
.superpowers/brainstorm/6600-1777284109/content/button-position.html View File

@@ -0,0 +1,72 @@
<h2>"去体验"按钮位置选择</h2>
<p class="subtitle">当前模型卡片结构示意,请选择按钮的最佳放置位置</p>

<div class="cards">
<div class="card" data-choice="bottom-center" onclick="toggleSelect(this)">
<div class="card-image">
<div style="background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); padding: 20px; border-radius: 8px; color: white;">
<div style="display:flex; align-items:center; gap:10px; margin-bottom:12px;">
<div style="width:36px;height:36px;border-radius:8px;background:rgba(255,255,255,0.2);display:flex;align-items:center;justify-content:center;font-size:18px;">GPT</div>
<span style="font-weight:600;font-size:16px;">gpt-4o</span>
<span style="margin-left:auto;font-size:12px;opacity:0.7;">📋</span>
</div>
<p style="font-size:13px;opacity:0.85;margin:0 0 12px 0;line-height:1.4;">OpenAI 旗舰多模态模型,支持文本和图像输入输出</p>
<div style="display:flex;gap:6px;margin-bottom:16px;">
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">按量计费</span>
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">多模态</span>
</div>
<button style="width:100%;padding:8px 0;background:rgba(255,255,255,0.25);border:1px solid rgba(255,255,255,0.4);border-radius:6px;color:white;font-size:13px;cursor:pointer;">去体验 →</button>
</div>
</div>
<div class="card-body">
<h3>方案 A:卡片底部居中</h3>
<p>按钮占据卡片底部整行,醒目且易于点击。适合卡片高度固定的场景。</p>
</div>
</div>

<div class="card" data-choice="footer-right" onclick="toggleSelect(this)">
<div class="card-image">
<div style="background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); padding: 20px; border-radius: 8px; color: white;">
<div style="display:flex; align-items:center; gap:10px; margin-bottom:12px;">
<div style="width:36px;height:36px;border-radius:8px;background:rgba(255,255,255,0.2);display:flex;align-items:center;justify-content:center;font-size:18px;">GPT</div>
<span style="font-weight:600;font-size:16px;">gpt-4o</span>
<span style="margin-left:auto;font-size:12px;opacity:0.7;">📋</span>
</div>
<p style="font-size:13px;opacity:0.85;margin:0 0 12px 0;line-height:1.4;">OpenAI 旗舰多模态模型,支持文本和图像输入输出</p>
<div style="display:flex;align-items:center;gap:6px;">
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">按量计费</span>
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">多模态</span>
<span style="margin-left:auto;"></span>
<button style="padding:5px 12px;background:rgba(255,255,255,0.25);border:1px solid rgba(255,255,255,0.4);border-radius:6px;color:white;font-size:12px;cursor:pointer;">去体验 →</button>
</div>
</div>
</div>
<div class="card-body">
<h3>方案 B:标签行右侧</h3>
<p>按钮与标签同行,紧凑不占额外空间。但可能在小屏幕上显得拥挤。</p>
</div>
</div>

<div class="card" data-choice="header-right" onclick="toggleSelect(this)">
<div class="card-image">
<div style="background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); padding: 20px; border-radius: 8px; color: white;">
<div style="display:flex; align-items:center; gap:10px; margin-bottom:12px;">
<div style="width:36px;height:36px;border-radius:8px;background:rgba(255,255,255,0.2);display:flex;align-items:center;justify-content:center;font-size:18px;">GPT</div>
<span style="font-weight:600;font-size:16px;">gpt-4o</span>
<span style="margin-left:auto;"></span>
<button style="padding:5px 12px;background:rgba(255,255,255,0.25);border:1px solid rgba(255,255,255,0.4);border-radius:6px;color:white;font-size:12px;cursor:pointer;">去体验 →</button>
<span style="font-size:12px;opacity:0.7;">📋</span>
</div>
<p style="font-size:13px;opacity:0.85;margin:0 0 12px 0;line-height:1.4;">OpenAI 旗舰多模态模型,支持文本和图像输入输出</p>
<div style="display:flex;gap:6px;">
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">按量计费</span>
<span style="background:rgba(255,255,255,0.2);padding:2px 8px;border-radius:4px;font-size:11px;">多模态</span>
</div>
</div>
</div>
<div class="card-body">
<h3>方案 C:标题行右侧</h3>
<p>按钮在模型名称旁边,最显眼的位置。但与复制按钮竞争空间。</p>
</div>
</div>
</div>

+ 3
- 0
.superpowers/brainstorm/6600-1777284109/content/waiting.html View File

@@ -0,0 +1,3 @@
<div style="display:flex;align-items:center;justify-content:center;min-height:60vh">
<p class="subtitle">Continuing in terminal...</p>
</div>

+ 1
- 0
.superpowers/brainstorm/6600-1777284109/state/server-stopped View File

@@ -0,0 +1 @@
{"reason":"idle timeout","timestamp":1777286270028}

+ 1
- 0
.superpowers/brainstorm/6600-1777284109/state/server.pid View File

@@ -0,0 +1 @@
6600

+ 6
- 5
Dockerfile View File

@@ -1,4 +1,4 @@
FROM oven/bun:latest AS builder
FROM swr.cn-north-4.myhuaweicloud.com/ddn-k8s/docker.io/oven/bun:latest AS builder

# 使用淘宝镜像加速
ENV BUN_INSTALL_REGISTRY=https://registry.npmmirror.com
@@ -6,10 +6,11 @@ ENV BUN_INSTALL_REGISTRY=https://registry.npmmirror.com
WORKDIR /build
COPY web/package.json .
COPY web/bun.lock .
RUN bun install
RUN --mount=type=cache,target=/root/.bun/install/cache bun install
COPY ./web .
COPY ./VERSION .
RUN DISABLE_ESLINT_PLUGIN='true' VITE_REACT_APP_VERSION=$(cat VERSION) bun run build
RUN cd widget && bun install && bun run build

FROM swr.cn-north-4.myhuaweicloud.com/ddn-k8s/docker.io/library/golang:1.26.1-alpine AS builder2
ENV GO111MODULE=on CGO_ENABLED=0
@@ -22,13 +23,13 @@ ENV GOOS=${TARGETOS:-linux} GOARCH=${TARGETARCH:-amd64}
WORKDIR /build

ADD go.mod go.sum ./
RUN go mod download
RUN --mount=type=cache,target=/go/pkg/mod go mod download

COPY . .
COPY --from=builder /build/dist ./web/dist
RUN go build -ldflags "-s -w -X 'github.com/QuantumNous/new-api/common.Version=$(cat VERSION)'" -o new-api
RUN --mount=type=cache,target=/root/.cache/go-build go build -ldflags "-s -w -X 'github.com/QuantumNous/new-api/common.Version=$(cat VERSION)'" -o new-api

FROM debian:bookworm-slim
FROM swr.cn-north-4.myhuaweicloud.com/ddn-k8s/docker.io/library/debian:bookworm-slim

# 使用阿里云镜像源
RUN sed -i 's/deb.debian.org/mirrors.aliyun.com/g' /etc/apt/sources.list.d/debian.sources


+ 123
- 0
alan-claudecode-429-upstream-ids.txt View File

@@ -0,0 +1,123 @@
# Alan 用户 ClaudeCode 令牌 429 错误 upstream_id 列表
# 用户: Alan (user_id=47), 令牌: ClaudeCode (token_id=14)
# 时间范围: 2026-05-19 23:47:19 ~ 2026-05-20 09:03:54
# 总计: 112 条 429 错误,全部来自 channel_id=1
# 模型分布: claude-opus-4.7 (90条), claude-haiku-4.5 (22条)
#
# 格式: upstream_id | model_name | time

# === claude-haiku-4.5-20251001 (22条) ===
8a4c62b4-1686-4f02-901e-7e7da6aac605 anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:12
259035f9-eb62-484a-9e7f-13a3885a6626 anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:17
72745e0f-6a18-4093-93ed-9d9d1149ef7a anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:20
2f6faaf6-5511-4e22-8ce2-044346c175dc anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:27
ec6e5677-8422-457a-870f-485a7c297790 anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:36
a2523ccf-7241-4a61-ae8a-1e1d1db20daf anthropic/claude-haiku-4.5-20251001 2026-05-20 09:00:50
429e40ab-058b-4999-be18-246d1d8a4926 anthropic/claude-haiku-4.5-20251001 2026-05-20 09:01:12
de03173d-c6e4-4327-8339-45cfa734066b anthropic/claude-haiku-4.5-20251001 2026-05-20 09:01:54
f01c38ab-a216-4da2-8823-e2107e54839d anthropic/claude-haiku-4.5-20251001 2026-05-20 09:02:33
9a7dbef1-dd8b-426b-a63f-cbbb5195682a anthropic/claude-haiku-4.5-20251001 2026-05-20 09:03:15
82a4bbf6-bc89-4617-8514-a4ac2cf5174f anthropic/claude-haiku-4.5-20251001 2026-05-20 09:03:54
40c0b99d-4c34-474c-a942-e2f42785646a anthropic/claude-haiku-4.5-20251001 2026-05-19 23:51:01
5fcf397a-1345-4174-9f92-a4f1e0607b1b anthropic/claude-haiku-4.5-20251001 2026-05-19 23:50:23
b22f2e85-07f7-4ffa-89de-2fa9ba5bb6a7 anthropic/claude-haiku-4.5-20251001 2026-05-19 23:49:47
c8f0f203-0253-41ae-b72e-80d8667dd017 anthropic/claude-haiku-4.5-20251001 2026-05-19 23:49:00
4ec13969-7f2a-41d6-83cf-b4f1829e8fae anthropic/claude-haiku-4.5-20251001 2026-05-19 23:48:22
aa3dd5ef-c73f-4527-b4b6-0ded7d445d92 anthropic/claude-haiku-4.5-20251001 2026-05-19 23:48:01
0435b06b-c437-4bc2-9d8a-b86a41110520 anthropic/claude-haiku-4.5-20251001 2026-05-19 23:47:47
b9c49d19-e532-4386-b932-e72c3964b10c anthropic/claude-haiku-4.5-20251001 2026-05-19 23:47:38
97bcded4-ce94-4ba5-833d-fd5f74675ec8 anthropic/claude-haiku-4.5-20251001 2026-05-19 23:47:31
38066d1a-0749-424b-99b6-3f9dafc344cc anthropic/claude-haiku-4.5-20251001 2026-05-19 23:47:25
0a3365ed-a54b-421d-8125-f384e6a18d5b anthropic/claude-haiku-4.5-20251001 2026-05-19 23:47:19

# === claude-opus-4.7 (90条) ===
8937e45a-090a-4f92-baa3-802ebdbadf9f anthropic/claude-opus-4.7 2026-05-20 06:19:10
2577f038-cb35-4b9c-9039-aa1a34ceced5 anthropic/claude-opus-4.7 2026-05-20 06:19:30
d6180a7a-4dda-489e-8f4e-ac32ecb8d960 anthropic/claude-opus-4.7 2026-05-20 06:19:48
9d105cfb-2bc8-4f54-80a8-03d8c02c0e4e anthropic/claude-opus-4.7 2026-05-20 06:20:10
fb1ce184-b582-4ce8-877e-4ff9b8c86c50 anthropic/claude-opus-4.7 2026-05-20 06:20:42
a3248024-75fb-431b-909b-6832f5f47800 anthropic/claude-opus-4.7 2026-05-20 06:21:13
5adca109-c90a-495c-8900-8168a0d0e1b3 anthropic/claude-opus-4.7 2026-05-20 06:21:54
3fc887f4-14a9-4ead-9473-de5f3befa5f4 anthropic/claude-opus-4.7 2026-05-20 06:22:45
9b385114-e385-4bf6-93f5-65d29d695ada anthropic/claude-opus-4.7 2026-05-20 06:23:37
12d1f1a4-963c-4710-aaa9-aedcec883d54 anthropic/claude-opus-4.7 2026-05-20 06:24:40
dd70e7e1-5658-48c5-ae41-bcc3648e0f3f anthropic/claude-opus-4.7 2026-05-20 06:25:41
97e7aaa9-c797-44d0-b241-39f4eb8ab314 anthropic/claude-opus-4.7 2026-05-20 06:32:03
dd3c3de5-bad7-4819-a736-0112589bc3a8 anthropic/claude-opus-4.7 2026-05-20 06:32:22
a599b1b9-0c60-4219-a7ab-0a0f595c425a anthropic/claude-opus-4.7 2026-05-20 06:32:40
d763ca6b-e38e-4d51-b6f2-d27d6816ad9a anthropic/claude-opus-4.7 2026-05-20 06:33:07
5131c19d-87f7-44f0-9fd7-70642c8fd1ec anthropic/claude-opus-4.7 2026-05-20 06:33:46
3515ded8-b04e-4af3-9509-6dab3e20a8a7 anthropic/claude-opus-4.7 2026-05-20 06:34:16
1ddcd66b-10be-4017-a9f1-2e678e95455a anthropic/claude-opus-4.7 2026-05-20 06:35:03
8425c233-7892-4dc9-b240-a4c7c9d7ece3 anthropic/claude-opus-4.7 2026-05-20 06:36:58
9690cf62-abbe-4e9d-ab8d-4e6a66ee384e anthropic/claude-opus-4.7 2026-05-20 06:39:16
c16b6d6d-52d9-4512-b864-6b7610bd200f anthropic/claude-opus-4.7 2026-05-20 06:43:08
0b1e715b-7dbf-499c-815a-837a50860631 anthropic/claude-opus-4.7 2026-05-20 06:44:19
7223d596-9dde-40e8-898a-c9f8a28c94ca anthropic/claude-opus-4.7 2026-05-20 06:44:51
cf39ec6b-c982-4c1d-9af7-8d0bc0a4a740 anthropic/claude-opus-4.7 2026-05-20 06:49:34
341e944e-0475-44e6-a316-d3ab69748d9d anthropic/claude-opus-4.7 2026-05-20 06:50:19
44b66588-71f0-45ba-add9-8326d89727a5 anthropic/claude-opus-4.7 2026-05-20 06:54:02
8fbf4a48-30cb-4928-95ac-d63570ac4e50 anthropic/claude-opus-4.7 2026-05-20 06:56:08
ba638c4b-fe4b-49e0-af48-61338523a16a anthropic/claude-opus-4.7 2026-05-20 06:57:18
2ed8faa7-c444-43e3-936c-e4254c12a9e7 anthropic/claude-opus-4.7 2026-05-20 06:57:48
8d3f86c8-a736-4f47-8b48-d6de8f0aed5b anthropic/claude-opus-4.7 2026-05-20 07:03:48
9b7858f5-e763-4ecf-836a-b66ec78b6e01 anthropic/claude-opus-4.7 2026-05-20 07:04:41
9b4eaf36-651e-4348-b9a0-d85d19e941e3 anthropic/claude-opus-4.7 2026-05-20 07:05:06
2147af82-839e-4e84-820f-abeb23c71289 anthropic/claude-opus-4.7 2026-05-20 07:05:32
34ed7bb3-99db-46f0-992f-3479223c168d anthropic/claude-opus-4.7 2026-05-20 07:06:00
050bd83b-10bb-4eef-8ba6-2ba80802bac8 anthropic/claude-opus-4.7 2026-05-20 07:06:44
2686785e-d2d5-4418-a8aa-f162de8924b4 anthropic/claude-opus-4.7 2026-05-20 07:07:21
12c8fd4f-d492-4fe1-8290-3e66523fcf56 anthropic/claude-opus-4.7 2026-05-20 07:07:59
f8231f97-4a77-49c9-aaf7-55c89ea170d6 anthropic/claude-opus-4.7 2026-05-20 07:08:56
bb4f79e7-7c35-413a-8d17-5968399133a6 anthropic/claude-opus-4.7 2026-05-20 07:10:07
a07961a7-41cd-41dd-9fce-49b244328a33 anthropic/claude-opus-4.7 2026-05-20 07:11:09
33d1d725-a065-49dd-a371-31983ccf58c8 anthropic/claude-opus-4.7 2026-05-20 07:12:06
8b5e2509-2512-4bf0-96ea-115e318143df anthropic/claude-opus-4.7 2026-05-20 07:15:13
51c579f4-80d1-43c1-897e-20c6dc69465d anthropic/claude-opus-4.7 2026-05-20 07:15:26
471666f2-2f73-46c5-bd76-13620b0fd3c9 anthropic/claude-opus-4.7 2026-05-20 07:15:47
abc8755a-4d56-4d36-8ca6-62ed68d2e979 anthropic/claude-opus-4.7 2026-05-20 07:16:05
dca16c52-2c99-43a9-8f9c-96d910beaba5 anthropic/claude-opus-4.7 2026-05-20 07:16:43
daa220d4-a311-4c33-bacd-72265ccf83cb anthropic/claude-opus-4.7 2026-05-20 07:17:05
3474a1ed-6f7a-48aa-8213-c201fca65c91 anthropic/claude-opus-4.7 2026-05-20 07:17:52
8f0f37f4-58dd-4b45-a0b4-53bddf35d1fe anthropic/claude-opus-4.7 2026-05-20 07:18:39
6cec0036-6c27-4e59-91a1-097c3aad49ab anthropic/claude-opus-4.7 2026-05-20 07:19:27
6b08f0cb-dfb7-4745-abb4-79f05b8d7b76 anthropic/claude-opus-4.7 2026-05-20 07:22:13
6408a6be-22a9-4a48-9b23-944e8e2ebdbe anthropic/claude-opus-4.7 2026-05-20 07:22:30
76f1a72b-4c91-4ca5-a889-f64c0e641353 anthropic/claude-opus-4.7 2026-05-20 07:22:47
8ecf0adc-c60f-434d-91e3-c3ec3f3e47f9 anthropic/claude-opus-4.7 2026-05-20 07:29:48
63283fcf-84cd-4778-b138-be9155ed2773 anthropic/claude-opus-4.7 2026-05-20 07:30:09
ae5c146f-127e-4cfa-a664-1496e25c9208 anthropic/claude-opus-4.7 2026-05-20 07:30:46
f3cf76fc-3741-4525-b7a3-dbcaf07651d4 anthropic/claude-opus-4.7 2026-05-20 07:31:09
075f9154-f857-4b59-b774-b85f4cb81a60 anthropic/claude-opus-4.7 2026-05-20 07:31:37
873a7295-d47c-4e3d-af8f-502d4a7e6a36 anthropic/claude-opus-4.7 2026-05-20 07:32:43
083f7117-4539-4bbc-8585-a0d520c8ad6b anthropic/claude-opus-4.7 2026-05-20 07:33:27
a12a28fe-50a0-4572-b23a-6f833939014c anthropic/claude-opus-4.7 2026-05-20 07:34:20
4c5b9328-e370-4ed9-835b-151701076447 anthropic/claude-opus-4.7 2026-05-20 07:35:20
3b9bb953-bdf3-4585-bc79-146a48d34a21 anthropic/claude-opus-4.7 2026-05-20 07:36:19
ffabd406-f4a8-41be-be59-099e30613fc1 anthropic/claude-opus-4.7 2026-05-20 07:37:15
61411ecc-5abb-4c47-bc15-055f27a51374 anthropic/claude-opus-4.7 2026-05-20 07:38:09
a898a6a4-3ee5-41d1-a1ef-2b47cd2d7501 anthropic/claude-opus-4.7 2026-05-20 07:38:23
f93ef5b5-4401-45e7-8ed8-59280b7b7ceb anthropic/claude-opus-4.7 2026-05-20 07:38:34
3f736901-0b44-4a7b-b1ff-9abea5688b9c anthropic/claude-opus-4.7 2026-05-20 07:38:44
09b55670-82ba-49e1-aa9f-376538a6d778 anthropic/claude-opus-4.7 2026-05-20 07:38:56
8efd0a1a-2cf3-44e6-976e-f99b6efb850d anthropic/claude-opus-4.7 2026-05-20 07:39:14
a2431acd-8f1c-4f54-ac48-0b9000629f21 anthropic/claude-opus-4.7 2026-05-20 07:39:40
e98d8348-7ed1-4677-90ee-51d180b863a8 anthropic/claude-opus-4.7 2026-05-20 07:40:43
10eb0a90-082e-47f2-a52c-a9cb64b518f7 anthropic/claude-opus-4.7 2026-05-20 08:43:09
0057a591-3727-402a-851f-b1e8610e3fca anthropic/claude-opus-4.7 2026-05-20 08:43:34
e3ded410-5c25-4160-a325-4b4f7a341c19 anthropic/claude-opus-4.7 2026-05-20 08:43:53
b3e197b9-0277-4e96-a533-b976effa9d64 anthropic/claude-opus-4.7 2026-05-20 08:44:34
241e074f-c084-405f-a46a-c4184b99b90a anthropic/claude-opus-4.7 2026-05-20 08:45:20
583b0e49-93f9-4138-84e1-e9fc2c4abaa3 anthropic/claude-opus-4.7 2026-05-20 08:45:51
f1381633-4701-4d2b-ab00-4e15ac1e4024 anthropic/claude-opus-4.7 2026-05-20 08:50:19
063bf354-28be-4e3e-9f5c-9f88c929acdc anthropic/claude-opus-4.7 2026-05-20 08:50:35
d44bc890-2756-4cd5-8ee4-60ba0543c8c4 anthropic/claude-opus-4.7 2026-05-20 08:51:16
628f9ce9-8cfe-4966-a33b-1b2df10b10d8 anthropic/claude-opus-4.7 2026-05-20 08:52:26
8884ae5d-d22a-4ea4-9247-818c4b6be648 anthropic/claude-opus-4.7 2026-05-20 08:52:44
b73a8fa5-4cb4-4481-90d8-33d485077a45 anthropic/claude-opus-4.7 2026-05-20 08:53:05
d3594203-fc5a-4858-8f8b-7f81f2ee52d5 anthropic/claude-opus-4.7 2026-05-20 08:54:28
71c44932-cf4a-423e-a8cc-14d18fbd0e42 anthropic/claude-opus-4.7 2026-05-20 08:55:20
183be15a-ab5d-470a-87b2-10b0032b5bc1 anthropic/claude-opus-4.7 2026-05-20 08:56:11
1e24fa4b-aeae-42a0-89c7-a08982edcbfb anthropic/claude-opus-4.7 2026-05-20 08:57:00
33191ba1-8252-4d81-beb8-0c39877408bc anthropic/claude-opus-4.7 2026-05-20 08:58:01
0293f22d-5042-42a3-837f-21df0cd1b53b anthropic/claude-opus-4.7 2026-05-19 23:56:18

BIN
atlas-edit-test.png View File

Before After
Width: 1  |  Height: 1  |  Size: 68 B

+ 23
- 0
atlas-gemini-image.json View File

@@ -0,0 +1,23 @@
{
"contents": {
"role": "USER",
"parts": [
{
"text": "Create a Tom and Jerry Poster."
}
]
},
"generationConfig": {
"responseModalities": [
"IMAGE"
],
"imageConfig": {
"aspectRatio": "16:9"
}
},
"safetySettings": {
"method": "PROBABILITY",
"category": "HARM_CATEGORY_DANGEROUS_CONTENT",
"threshold": "BLOCK_MEDIUM_AND_ABOVE"
}
}

BIN
atlascloud-local-test.png View File

Before After
Width: 1024  |  Height: 1024  |  Size: 860 KiB

+ 4
- 0
common/api_type.go View File

@@ -53,6 +53,10 @@ func ChannelType2APIType(channelType int) (int, bool) {
apiType = constant.APITypeMokaAI
case constant.ChannelTypeVolcEngine:
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeDoubaoVideoCompatibleAiping:
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeDoubaoVideoCompatibleTianyiYun:
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeBaiduV2:
apiType = constant.APITypeBaiduV2
case constant.ChannelTypeOpenRouter:


+ 30
- 0
common/captcha.go View File

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

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

var captchaStore = base64Captcha.DefaultMemStore

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

var captchaInstance = base64Captcha.NewCaptcha(captchaDriver, captchaStore)

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

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

+ 58
- 0
common/codex_credential.go View File

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

import (
"errors"
"strings"
)

type CodexCredentialMode string

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

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

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

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

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

return &credential, nil
}

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

+ 10
- 0
common/constants.go View File

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

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

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

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


+ 1
- 0
common/endpoint_defaults.go View File

@@ -25,6 +25,7 @@ var defaultEndpointInfoMap = map[constant.EndpointType]EndpointInfo{
constant.EndpointTypeJinaRerank: {Path: "/v1/rerank", Method: "POST"},
constant.EndpointTypeImageGeneration: {Path: "/v1/images/generations", Method: "POST"},
constant.EndpointTypeEmbeddings: {Path: "/v1/embeddings", Method: "POST"},
constant.EndpointTypeDoubaoVideo: {Path: "/api/v3/contents/generations/tasks", Method: "POST"},
}

// GetDefaultEndpointInfo 返回指定端点类型的默认信息以及是否存在


+ 2
- 0
common/endpoint_type.go View File

@@ -30,6 +30,8 @@ func GetEndpointTypesByChannelType(channelType int, modelName string) []constant
endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAI, constant.EndpointTypeOpenAIResponse}
case constant.ChannelTypeSora:
endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAIVideo}
case constant.ChannelTypeDoubaoVideo, constant.ChannelTypeDoubaoVideoCompatibleAiping, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun:
endpointTypes = []constant.EndpointType{constant.EndpointTypeDoubaoVideo}
default:
if IsOpenAIResponseOnlyModel(modelName) {
endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAIResponse}


+ 24
- 4
common/limiter/limiter.go View File

@@ -13,9 +13,13 @@ import (
//go:embed lua/rate_limit.lua
var rateLimitScript string

//go:embed lua/rate_refund.lua
var rateRefundScript string

type RedisLimiter struct {
client *redis.Client
limitScriptSHA string
client *redis.Client
limitScriptSHA string
refundScriptSHA string
}

var (
@@ -30,9 +34,14 @@ func New(ctx context.Context, r *redis.Client) *RedisLimiter {
if err != nil {
common.SysLog(fmt.Sprintf("Failed to load rate limit script: %v", err))
}
refundSHA, err := r.ScriptLoad(ctx, rateRefundScript).Result()
if err != nil {
common.SysLog(fmt.Sprintf("Failed to load rate refund script: %v", err))
}
instance = &RedisLimiter{
client: r,
limitScriptSHA: limitSHA,
client: r,
limitScriptSHA: limitSHA,
refundScriptSHA: refundSHA,
}
})

@@ -68,6 +77,17 @@ func (rl *RedisLimiter) Allow(ctx context.Context, key string, opts ...Option) (
return result == 1, nil
}

func (rl *RedisLimiter) Refund(ctx context.Context, key string, requested, capacity int64) error {
_, err := rl.client.EvalSha(
ctx,
rl.refundScriptSHA,
[]string{key},
requested,
capacity,
).Int()
return err
}

// Config 配置选项模式
type Config struct {
Capacity int64


+ 18
- 0
common/limiter/lua/rate_refund.lua View File

@@ -0,0 +1,18 @@
-- 令牌桶退款脚本
-- KEYS[1]: 限流器唯一标识
-- ARGV[1]: 退还的令牌数
-- ARGV[2]: 桶容量(上限)
local key = KEYS[1]
local refund = tonumber(ARGV[1])
local capacity = tonumber(ARGV[2])

local bucket = redis.call('HMGET', key, 'tokens', 'last_time')
local tokens = tonumber(bucket[1])

if not tokens then
return 0
end

tokens = math.min(capacity, tokens + refund)
redis.call('HSET', key, 'tokens', tokens)
return 1

+ 64
- 0
common/metrics/collector.go View File

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

import (
"fmt"

"github.com/prometheus/client_golang/prometheus"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
)

// ChannelCollector 从数据库查询渠道状态并暴露为 Prometheus Gauge
type ChannelCollector struct {
channelStatus *prometheus.Desc
}

// NewChannelCollector 创建渠道状态 Collector
func NewChannelCollector() *ChannelCollector {
return &ChannelCollector{
channelStatus: prometheus.NewDesc(
"newapi_channel_status",
"Channel enabled status (1=enabled, 0=disabled)",
[]string{"channel", "channel_type"},
nil,
),
}
}

func (c *ChannelCollector) Describe(ch chan<- *prometheus.Desc) {
ch <- c.channelStatus
}

func (c *ChannelCollector) Collect(ch chan<- prometheus.Metric) {
if model.DB == nil {
return
}

var channels []struct {
Id int
Type int
Status int
}

err := model.DB.Table("channels").Select("id, type, status").Find(&channels).Error
if err != nil {
return
}

for _, ch2 := range channels {
value := float64(0)
if ch2.Status == common.ChannelStatusEnabled {
value = 1
}
m, err := prometheus.NewConstMetric(
c.channelStatus,
prometheus.GaugeValue,
value,
fmt.Sprintf("%d", ch2.Id),
fmt.Sprintf("%d", ch2.Type),
)
if err == nil {
ch <- m
}
}
}

+ 93
- 0
common/metrics/collector_test.go View File

@@ -0,0 +1,93 @@
package metrics

import (
"testing"

"github.com/glebarez/sqlite"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"

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

func setupCollectorTestDB(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.Channel{}))

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

return db
}

func TestChannelCollector_CollectsChannelStatus(t *testing.T) {
db := setupCollectorTestDB(t)

ch1 := model.Channel{Id: 1, Name: "openai-1", Type: 1, Status: common.ChannelStatusEnabled}
ch2 := model.Channel{Id: 2, Name: "claude-1", Type: 14, Status: common.ChannelStatusEnabled}
ch3 := model.Channel{Id: 3, Name: "gemini-1", Type: 24, Status: 3}
require.NoError(t, db.Create(&ch1).Error)
require.NoError(t, db.Create(&ch2).Error)
require.NoError(t, db.Create(&ch3).Error)

collector := NewChannelCollector()
reg := prometheus.NewRegistry()
reg.MustRegister(collector)

mfs, err := reg.Gather()
require.NoError(t, err)

var found bool
for _, mf := range mfs {
if mf.GetName() == "newapi_channel_status" {
found = true
metrics := mf.GetMetric()
assert.Len(t, metrics, 3)

statusMap := map[string]float64{}
for _, m := range metrics {
chId := ""
for _, label := range m.GetLabel() {
if label.GetName() == "channel" {
chId = label.GetValue()
}
}
statusMap[chId] = m.GetGauge().GetValue()
}
assert.Equal(t, float64(1), statusMap["1"])
assert.Equal(t, float64(1), statusMap["2"])
assert.Equal(t, float64(0), statusMap["3"])
}
}
assert.True(t, found, "newapi_channel_status metric not found")
}

func TestChannelCollector_EmptyDB(t *testing.T) {
setupCollectorTestDB(t)

collector := NewChannelCollector()
reg := prometheus.NewRegistry()
reg.MustRegister(collector)

mfs, err := reg.Gather()
require.NoError(t, err)

for _, mf := range mfs {
assert.NotEqual(t, "newapi_channel_status", mf.GetName(), "should not have channel status with empty DB")
}
}

+ 53
- 0
common/metrics/errorclassify.go View File

@@ -0,0 +1,53 @@
package metrics

import (
"strconv"
"strings"
)

// ClassifyError 根据状态码和错误消息分类错误
func ClassifyError(statusCode int, errMsg string) (string, string) {
// 超时类
if strings.Contains(errMsg, "context deadline exceeded") {
return "timeout", "context_deadline_exceeded"
}
if strings.Contains(errMsg, "context canceled") {
return "timeout", "context_canceled"
}

// 配额/计费类
if statusCode == 402 {
return "billing", "insufficient_quota"
}
if statusCode == 403 && strings.Contains(strings.ToLower(errMsg), "quota") {
return "billing", "quota_exceeded"
}

// 认证类
if statusCode == 401 {
return "auth", "401"
}
if statusCode == 403 && !strings.Contains(strings.ToLower(errMsg), "quota") {
return "auth", "403"
}

// 速率限制
if statusCode == 429 {
return "rate_limit", "429"
}

// 上游错误(5xx)
if statusCode >= 500 {
return "upstream", strconv.Itoa(statusCode)
}
if statusCode >= 400 {
return "other", strconv.Itoa(statusCode)
}

// 无状态码的未知错误
if statusCode == 0 {
return "other", "unknown"
}

return "other", strconv.Itoa(statusCode)
}

+ 85
- 0
common/metrics/errorclassify_test.go View File

@@ -0,0 +1,85 @@
package metrics

import (
"testing"

"github.com/stretchr/testify/assert"
)

func TestClassifyError_RateLimitEmptyMsg(t *testing.T) {
errorType, errorCode := ClassifyError(429, "")
assert.Equal(t, "rate_limit", errorType)
assert.Equal(t, "429", errorCode)
}

func TestClassifyError_UpstreamHTTP502(t *testing.T) {
errorType, errorCode := ClassifyError(502, "")
assert.Equal(t, "upstream", errorType)
assert.Equal(t, "502", errorCode)
}

func TestClassifyError_UpstreamHTTP500(t *testing.T) {
errorType, errorCode := ClassifyError(500, "")
assert.Equal(t, "upstream", errorType)
assert.Equal(t, "500", errorCode)
}

func TestClassifyError_Timeout(t *testing.T) {
errorType, errorCode := ClassifyError(0, "context deadline exceeded")
assert.Equal(t, "timeout", errorType)
assert.Equal(t, "context_deadline_exceeded", errorCode)
}

func TestClassifyError_TimeoutCanceled(t *testing.T) {
errorType, errorCode := ClassifyError(0, "context canceled")
assert.Equal(t, "timeout", errorType)
assert.Equal(t, "context_canceled", errorCode)
}

func TestClassifyError_AuthInvalidToken(t *testing.T) {
errorType, errorCode := ClassifyError(401, "invalid api key")
assert.Equal(t, "auth", errorType)
assert.Equal(t, "401", errorCode)
}

func TestClassifyError_AuthForbidden(t *testing.T) {
errorType, errorCode := ClassifyError(403, "")
assert.Equal(t, "auth", errorType)
assert.Equal(t, "403", errorCode)
}

func TestClassifyError_RateLimit(t *testing.T) {
errorType, errorCode := ClassifyError(429, "rate limit exceeded")
assert.Equal(t, "rate_limit", errorType)
assert.Equal(t, "429", errorCode)
}

func TestClassifyError_BillingInsufficientQuota(t *testing.T) {
errorType, errorCode := ClassifyError(402, "insufficient quota")
assert.Equal(t, "billing", errorType)
assert.Equal(t, "insufficient_quota", errorCode)
}

func TestClassifyError_BillingQuotaExceeded(t *testing.T) {
errorType, errorCode := ClassifyError(403, "quota exceeded")
assert.Equal(t, "billing", errorType)
assert.Equal(t, "quota_exceeded", errorCode)
}

func TestClassifyError_Other(t *testing.T) {
errorType, errorCode := ClassifyError(0, "something unexpected")
assert.Equal(t, "other", errorType)
assert.Equal(t, "unknown", errorCode)
}

func TestClassifyError_BadRequest(t *testing.T) {
errorType, errorCode := ClassifyError(400, "")
assert.Equal(t, "other", errorType)
assert.Equal(t, "400", errorCode)
}

func TestClassifyError_EmptyBoth(t *testing.T) {
errorType, errorCode := ClassifyError(0, "")
assert.Equal(t, "other", errorType)
assert.Equal(t, "unknown", errorCode)
}

+ 119
- 0
common/metrics/metrics.go View File

@@ -0,0 +1,119 @@
package metrics

import (
"net/http"
"os"
"sync"

"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promhttp"
)

var (
globalMetrics *Metrics
globalOnce sync.Once
)

// Metrics 包含所有 Prometheus 指标
type Metrics struct {
RequestsTotal *prometheus.CounterVec
RequestErrorsTotal *prometheus.CounterVec
RequestRetriesTotal *prometheus.CounterVec
TokensTotal *prometheus.CounterVec
QuotaConsumedTotal *prometheus.CounterVec

RequestDuration *prometheus.HistogramVec
UpstreamDuration *prometheus.HistogramVec
FirstTokenDuration *prometheus.HistogramVec

ActiveRequests *prometheus.GaugeVec
}

// NewMetrics 创建并注册所有指标到给定 Registry
func NewMetrics(reg prometheus.Registerer) *Metrics {
m := &Metrics{
RequestsTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Name: "newapi_requests_total",
Help: "Total number of API requests",
}, []string{"channel", "model", "code", "is_stream", "relay_mode"}),
RequestErrorsTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Name: "newapi_request_errors_total",
Help: "Total number of request errors",
}, []string{"channel", "model", "error_type", "error_code"}),
RequestRetriesTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Name: "newapi_request_retries_total",
Help: "Total number of request retries",
}, []string{"channel", "model"}),
TokensTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Name: "newapi_tokens_total",
Help: "Total tokens used",
}, []string{"channel", "model", "type"}),
QuotaConsumedTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Name: "newapi_quota_consumed_total",
Help: "Total quota consumed",
}, []string{"channel", "model", "billing_source"}),

RequestDuration: prometheus.NewHistogramVec(prometheus.HistogramOpts{
Name: "newapi_request_duration_seconds",
Help: "End-to-end request duration in seconds",
Buckets: []float64{0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30, 60, 120},
}, []string{"channel", "model", "is_stream", "relay_mode"}),
UpstreamDuration: prometheus.NewHistogramVec(prometheus.HistogramOpts{
Name: "newapi_upstream_duration_seconds",
Help: "Upstream response duration in seconds",
Buckets: []float64{0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30, 60, 120},
}, []string{"channel", "model"}),
FirstTokenDuration: prometheus.NewHistogramVec(prometheus.HistogramOpts{
Name: "newapi_first_token_duration_seconds",
Help: "Time to first token in seconds (streaming only)",
Buckets: []float64{0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10},
}, []string{"channel", "model"}),

ActiveRequests: prometheus.NewGaugeVec(prometheus.GaugeOpts{
Name: "newapi_active_requests",
Help: "Current number of active requests",
}, []string{"model"}),
}

reg.MustRegister(
m.RequestsTotal,
m.RequestErrorsTotal,
m.RequestRetriesTotal,
m.TokensTotal,
m.QuotaConsumedTotal,
m.RequestDuration,
m.UpstreamDuration,
m.FirstTokenDuration,
m.ActiveRequests,
)

return m
}

// GetMetrics 返回全局 Metrics 实例
func GetMetrics() *Metrics {
globalOnce.Do(func() {
if os.Getenv("METRICS_ENABLED") == "false" {
globalMetrics = NewMetrics(prometheus.NewPedanticRegistry())
return
}
globalMetrics = NewMetrics(prometheus.DefaultRegisterer)
prometheus.MustRegister(NewChannelCollector())
})
return globalMetrics
}

// IsEnabled 返回 metrics 是否启用
func IsEnabled() bool {
return os.Getenv("METRICS_ENABLED") != "false"
}

// NewMetricsHandler 返回默认 registry 的 /metrics HTTP handler
func NewMetricsHandler() http.Handler {
return promhttp.Handler()
}

// NewMetricsHandlerWithRegistry 返回指定 registry 的 /metrics HTTP handler
func NewMetricsHandlerWithRegistry(reg *prometheus.Registry) http.Handler {
return promhttp.HandlerFor(reg, promhttp.HandlerOpts{})
}

+ 242
- 0
common/metrics/metrics_test.go View File

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

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

"github.com/gin-gonic/gin"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/testutil"
"github.com/stretchr/testify/assert"
)

func newTestRegistry() *prometheus.Registry {
return prometheus.NewRegistry()
}

func TestRequestsTotalCounter(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.RequestsTotal.WithLabelValues("1", "gpt-4", "200", "true", "chat").Inc()
m.RequestsTotal.WithLabelValues("1", "gpt-4", "200", "true", "chat").Inc()
m.RequestsTotal.WithLabelValues("2", "claude-3", "429", "false", "chat").Inc()

assert.Equal(t, float64(2), testutil.ToFloat64(m.RequestsTotal.WithLabelValues("1", "gpt-4", "200", "true", "chat")))
assert.Equal(t, float64(1), testutil.ToFloat64(m.RequestsTotal.WithLabelValues("2", "claude-3", "429", "false", "chat")))
}

func TestRequestErrorsTotalCounter(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.RequestErrorsTotal.WithLabelValues("1", "gpt-4", "upstream", "429").Inc()
m.RequestErrorsTotal.WithLabelValues("1", "gpt-4", "upstream", "429").Inc()
m.RequestErrorsTotal.WithLabelValues("2", "claude-3", "timeout", "context_deadline_exceeded").Inc()

assert.Equal(t, float64(2), testutil.ToFloat64(m.RequestErrorsTotal.WithLabelValues("1", "gpt-4", "upstream", "429")))
assert.Equal(t, float64(1), testutil.ToFloat64(m.RequestErrorsTotal.WithLabelValues("2", "claude-3", "timeout", "context_deadline_exceeded")))
}

func TestRequestRetriesTotalCounter(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.RequestRetriesTotal.WithLabelValues("1", "gpt-4").Inc()
m.RequestRetriesTotal.WithLabelValues("1", "gpt-4").Inc()

assert.Equal(t, float64(2), testutil.ToFloat64(m.RequestRetriesTotal.WithLabelValues("1", "gpt-4")))
}

func TestTokensTotalCounter(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.TokensTotal.WithLabelValues("1", "gpt-4", "prompt").Add(100)
m.TokensTotal.WithLabelValues("1", "gpt-4", "completion").Add(50)
m.TokensTotal.WithLabelValues("1", "gpt-4", "cache").Add(20)

assert.Equal(t, float64(100), testutil.ToFloat64(m.TokensTotal.WithLabelValues("1", "gpt-4", "prompt")))
assert.Equal(t, float64(50), testutil.ToFloat64(m.TokensTotal.WithLabelValues("1", "gpt-4", "completion")))
assert.Equal(t, float64(20), testutil.ToFloat64(m.TokensTotal.WithLabelValues("1", "gpt-4", "cache")))
}

func TestQuotaConsumedTotalCounter(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.QuotaConsumedTotal.WithLabelValues("1", "gpt-4", "wallet").Add(5000)

assert.Equal(t, float64(5000), testutil.ToFloat64(m.QuotaConsumedTotal.WithLabelValues("1", "gpt-4", "wallet")))
}

func TestRequestDurationHistogram(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.RequestDuration.WithLabelValues("1", "gpt-4", "true", "chat").Observe(1.5)
m.RequestDuration.WithLabelValues("1", "gpt-4", "true", "chat").Observe(2.5)

// 验证 histogram metric family 被注册且有数据
assert.Equal(t, 1, testutil.CollectAndCount(m.RequestDuration))
}

func TestUpstreamDurationHistogram(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.UpstreamDuration.WithLabelValues("1", "gpt-4").Observe(0.5)

assert.Equal(t, 1, testutil.CollectAndCount(m.UpstreamDuration))
}

func TestFirstTokenDurationHistogram(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.FirstTokenDuration.WithLabelValues("1", "gpt-4").Observe(0.1)

assert.Equal(t, 1, testutil.CollectAndCount(m.FirstTokenDuration))
}

func TestActiveRequestsGauge(t *testing.T) {
reg := newTestRegistry()
m := NewMetrics(reg)

m.ActiveRequests.WithLabelValues("gpt-4").Inc()
m.ActiveRequests.WithLabelValues("gpt-4").Inc()
m.ActiveRequests.WithLabelValues("claude-3").Inc()

assert.Equal(t, float64(2), testutil.ToFloat64(m.ActiveRequests.WithLabelValues("gpt-4")))
assert.Equal(t, float64(1), testutil.ToFloat64(m.ActiveRequests.WithLabelValues("claude-3")))

m.ActiveRequests.WithLabelValues("gpt-4").Dec()
assert.Equal(t, float64(1), testutil.ToFloat64(m.ActiveRequests.WithLabelValues("gpt-4")))
}

func TestMetricsDisabledIsNoOp(t *testing.T) {
// 当 METRICS_ENABLED=false 时,使用独立 registry(Discard)
reg := prometheus.NewPedanticRegistry()
m := NewMetrics(reg)
// 这些调用不应该 panic
m.RequestsTotal.WithLabelValues("1", "gpt-4", "200", "true", "chat").Inc()
m.ActiveRequests.WithLabelValues("gpt-4").Inc()
m.RequestDuration.WithLabelValues("1", "gpt-4", "true", "chat").Observe(1.0)
}

func TestMetricsMiddleware_ActiveRequests(t *testing.T) {
gin.SetMode(gin.TestMode)
reg := newTestRegistry()
m := NewMetrics(reg)

r := gin.New()
r.Use(MetricsMiddleware(m))

blockCh := make(chan struct{})
r.GET("/test", func(c *gin.Context) {
<-blockCh
c.Status(200)
})

go func() {
req := httptest.NewRequest("GET", "/test", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
}()
go func() {
req := httptest.NewRequest("GET", "/test", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
}()

time.Sleep(100 * time.Millisecond)

assert.Equal(t, float64(2), testutil.ToFloat64(m.ActiveRequests.WithLabelValues("")))

close(blockCh)
time.Sleep(100 * time.Millisecond)

assert.Equal(t, float64(0), testutil.ToFloat64(m.ActiveRequests.WithLabelValues("")))
}

func TestMetricsMiddleware_RequestsTotalByCode(t *testing.T) {
gin.SetMode(gin.TestMode)
reg := newTestRegistry()
m := NewMetrics(reg)

r := gin.New()
r.Use(MetricsMiddleware(m))
r.GET("/200", func(c *gin.Context) { c.Status(200) })
r.GET("/400", func(c *gin.Context) { c.Status(400) })
r.GET("/500", func(c *gin.Context) { c.Status(500) })

for i := 0; i < 3; i++ {
req := httptest.NewRequest("GET", "/200", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
}
req := httptest.NewRequest("GET", "/400", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
req = httptest.NewRequest("GET", "/500", nil)
w = httptest.NewRecorder()
r.ServeHTTP(w, req)

assert.Equal(t, float64(3), testutil.ToFloat64(m.RequestsTotal.WithLabelValues("", "", "200", "", "")))
assert.Equal(t, float64(1), testutil.ToFloat64(m.RequestsTotal.WithLabelValues("", "", "400", "", "")))
assert.Equal(t, float64(1), testutil.ToFloat64(m.RequestsTotal.WithLabelValues("", "", "500", "", "")))
}

func TestMetricsMiddleware_RequestDuration(t *testing.T) {
gin.SetMode(gin.TestMode)
reg := newTestRegistry()
m := NewMetrics(reg)

r := gin.New()
r.Use(MetricsMiddleware(m))
r.GET("/slow", func(c *gin.Context) {
time.Sleep(100 * time.Millisecond)
c.Status(200)
})

req := httptest.NewRequest("GET", "/slow", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)

assert.Equal(t, 1, testutil.CollectAndCount(m.RequestDuration))
}

func TestMetricsEndpoint_Format(t *testing.T) {
handler := NewMetricsHandler()

req := httptest.NewRequest("GET", "/metrics", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)

assert.Equal(t, 200, w.Code)
assert.Contains(t, w.Header().Get("Content-Type"), "text/plain")
}

func TestMetricsEndpoint_WithData(t *testing.T) {
reg := newTestRegistry()
_ = NewMetrics(reg)

handler := NewMetricsHandlerWithRegistry(reg)

req := httptest.NewRequest("GET", "/metrics", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)

assert.Equal(t, 200, w.Code)
}

func TestMetricsEndpoint_EmptyNoPanic(t *testing.T) {
handler := NewMetricsHandler()

req := httptest.NewRequest("GET", "/metrics", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)

assert.Equal(t, 200, w.Code)
}

+ 28
- 0
common/metrics/middleware.go View File

@@ -0,0 +1,28 @@
package metrics

import (
"strconv"
"time"

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

// MetricsMiddleware 是一个 Gin 中间件,采集活跃请求数、请求延迟和状态码。
func MetricsMiddleware(m *Metrics) gin.HandlerFunc {
return func(c *gin.Context) {
start := time.Now()

m.ActiveRequests.WithLabelValues("").Inc()
defer func() {
m.ActiveRequests.WithLabelValues("").Dec()

elapsed := time.Since(start).Seconds()
status := strconv.Itoa(c.Writer.Status())

m.RequestsTotal.WithLabelValues("", "", status, "", "").Inc()
m.RequestDuration.WithLabelValues("", "", "", "").Observe(elapsed)
}()

c.Next()
}
}

+ 12
- 0
common/rate-limit.go View File

@@ -68,3 +68,15 @@ func (l *InMemoryRateLimiter) Request(key string, maxRequestNum int, duration in
}
return true
}

// Refund 移除 key 最近一次请求记录(用于请求失败时退还配额)
func (l *InMemoryRateLimiter) Refund(key string) bool {
l.mutex.Lock()
defer l.mutex.Unlock()
queue, ok := l.store[key]
if !ok || len(*queue) == 0 {
return false
}
*queue = (*queue)[:len(*queue)-1]
return true
}

+ 18
- 0
common/str.go View File

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

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

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


+ 21
- 0
common/tianyiyun_channel_test.go View File

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

import (
"testing"

"github.com/QuantumNous/new-api/constant"
"github.com/stretchr/testify/require"
)

func TestTianyiYunChannelUsesVolcEngineAPIType(t *testing.T) {
apiType, ok := ChannelType2APIType(constant.ChannelTypeDoubaoVideoCompatibleTianyiYun)

require.True(t, ok)
require.Equal(t, constant.APITypeVolcEngine, apiType)
}

func TestTianyiYunChannelUsesDoubaoVideoEndpointType(t *testing.T) {
got := GetEndpointTypesByChannelType(constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "cdance2.0-0611")

require.Equal(t, []constant.EndpointType{constant.EndpointTypeDoubaoVideo}, got)
}

+ 4
- 0
common/utils.go View File

@@ -276,6 +276,10 @@ func Max(a int, b int) int {
}

func MessageWithRequestId(message string, id string) string {
id = strings.TrimSpace(id)
if id == "" || strings.Contains(message, "request id:") {
return message
}
return fmt.Sprintf("%s (request id: %s)", message, id)
}



+ 27
- 0
common/utils_test.go View File

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

import "testing"

func TestMessageWithRequestIdAppendsRequestId(t *testing.T) {
got := MessageWithRequestId("upstream failed", "req-123")
want := "upstream failed (request id: req-123)"
if got != want {
t.Fatalf("MessageWithRequestId() = %q, want %q", got, want)
}
}

func TestMessageWithRequestIdSkipsEmptyRequestId(t *testing.T) {
got := MessageWithRequestId("upstream failed", " ")
want := "upstream failed"
if got != want {
t.Fatalf("MessageWithRequestId() = %q, want %q", got, want)
}
}

func TestMessageWithRequestIdDoesNotAppendTwice(t *testing.T) {
message := "upstream failed (request id: req-123)"
got := MessageWithRequestId(message, "req-123")
if got != message {
t.Fatalf("MessageWithRequestId() = %q, want %q", got, message)
}
}

+ 118
- 109
constant/channel.go View File

@@ -1,61 +1,64 @@
package constant

const (
ChannelTypeUnknown = 0
ChannelTypeOpenAI = 1
ChannelTypeMidjourney = 2
ChannelTypeAzure = 3
ChannelTypeOllama = 4
ChannelTypeMidjourneyPlus = 5
ChannelTypeOpenAIMax = 6
ChannelTypeOhMyGPT = 7
ChannelTypeCustom = 8
ChannelTypeAILS = 9
ChannelTypeAIProxy = 10
ChannelTypePaLM = 11
ChannelTypeAPI2GPT = 12
ChannelTypeAIGC2D = 13
ChannelTypeAnthropic = 14
ChannelTypeBaidu = 15
ChannelTypeZhipu = 16
ChannelTypeAli = 17
ChannelTypeXunfei = 18
ChannelType360 = 19
ChannelTypeOpenRouter = 20
ChannelTypeAIProxyLibrary = 21
ChannelTypeFastGPT = 22
ChannelTypeTencent = 23
ChannelTypeGemini = 24
ChannelTypeMoonshot = 25
ChannelTypeZhipu_v4 = 26
ChannelTypePerplexity = 27
ChannelTypeLingYiWanWu = 31
ChannelTypeAws = 33
ChannelTypeCohere = 34
ChannelTypeMiniMax = 35
ChannelTypeSunoAPI = 36
ChannelTypeDify = 37
ChannelTypeJina = 38
ChannelCloudflare = 39
ChannelTypeSiliconFlow = 40
ChannelTypeVertexAi = 41
ChannelTypeMistral = 42
ChannelTypeDeepSeek = 43
ChannelTypeMokaAI = 44
ChannelTypeVolcEngine = 45
ChannelTypeBaiduV2 = 46
ChannelTypeXinference = 47
ChannelTypeXai = 48
ChannelTypeCoze = 49
ChannelTypeKling = 50
ChannelTypeJimeng = 51
ChannelTypeVidu = 52
ChannelTypeSubmodel = 53
ChannelTypeDoubaoVideo = 54
ChannelTypeSora = 55
ChannelTypeReplicate = 56
ChannelTypeCodex = 57
ChannelTypeDummy // this one is only for count, do not add any channel after this
ChannelTypeUnknown = 0
ChannelTypeOpenAI = 1
ChannelTypeMidjourney = 2
ChannelTypeAzure = 3
ChannelTypeOllama = 4
ChannelTypeMidjourneyPlus = 5
ChannelTypeOpenAIMax = 6
ChannelTypeOhMyGPT = 7
ChannelTypeCustom = 8
ChannelTypeAILS = 9
ChannelTypeAIProxy = 10
ChannelTypePaLM = 11
ChannelTypeAPI2GPT = 12
ChannelTypeAIGC2D = 13
ChannelTypeAnthropic = 14
ChannelTypeBaidu = 15
ChannelTypeZhipu = 16
ChannelTypeAli = 17
ChannelTypeXunfei = 18
ChannelType360 = 19
ChannelTypeOpenRouter = 20
ChannelTypeAIProxyLibrary = 21
ChannelTypeFastGPT = 22
ChannelTypeTencent = 23
ChannelTypeGemini = 24
ChannelTypeMoonshot = 25
ChannelTypeZhipu_v4 = 26
ChannelTypePerplexity = 27
ChannelTypeLingYiWanWu = 31
ChannelTypeAws = 33
ChannelTypeCohere = 34
ChannelTypeMiniMax = 35
ChannelTypeSunoAPI = 36
ChannelTypeDify = 37
ChannelTypeJina = 38
ChannelCloudflare = 39
ChannelTypeSiliconFlow = 40
ChannelTypeVertexAi = 41
ChannelTypeMistral = 42
ChannelTypeDeepSeek = 43
ChannelTypeMokaAI = 44
ChannelTypeVolcEngine = 45
ChannelTypeBaiduV2 = 46
ChannelTypeXinference = 47
ChannelTypeXai = 48
ChannelTypeCoze = 49
ChannelTypeKling = 50
ChannelTypeJimeng = 51
ChannelTypeVidu = 52
ChannelTypeSubmodel = 53
ChannelTypeDoubaoVideo = 54
ChannelTypeSora = 55
ChannelTypeReplicate = 56
ChannelTypeCodex = 57
ChannelTypeDoubaoVideoCompatibleAiping = 58
ChannelTypeKlingAiping = 59
ChannelTypeDoubaoVideoCompatibleTianyiYun = 60
ChannelTypeDummy // this one is only for count, do not add any channel after this

)

@@ -118,63 +121,69 @@ var ChannelBaseURLs = []string{
"https://api.openai.com", //55
"https://api.replicate.com", //56
"https://chatgpt.com", //57
"", //58
"https://aiping.cn/api", //59
"https://ai.ctaigw.cn", //60
}

var ChannelTypeNames = map[int]string{
ChannelTypeUnknown: "Unknown",
ChannelTypeOpenAI: "OpenAI",
ChannelTypeMidjourney: "Midjourney",
ChannelTypeAzure: "Azure",
ChannelTypeOllama: "Ollama",
ChannelTypeMidjourneyPlus: "MidjourneyPlus",
ChannelTypeOpenAIMax: "OpenAIMax",
ChannelTypeOhMyGPT: "OhMyGPT",
ChannelTypeCustom: "Custom",
ChannelTypeAILS: "AILS",
ChannelTypeAIProxy: "AIProxy",
ChannelTypePaLM: "PaLM",
ChannelTypeAPI2GPT: "API2GPT",
ChannelTypeAIGC2D: "AIGC2D",
ChannelTypeAnthropic: "Anthropic",
ChannelTypeBaidu: "Baidu",
ChannelTypeZhipu: "Zhipu",
ChannelTypeAli: "Ali",
ChannelTypeXunfei: "Xunfei",
ChannelType360: "360",
ChannelTypeOpenRouter: "OpenRouter",
ChannelTypeAIProxyLibrary: "AIProxyLibrary",
ChannelTypeFastGPT: "FastGPT",
ChannelTypeTencent: "Tencent",
ChannelTypeGemini: "Gemini",
ChannelTypeMoonshot: "Moonshot",
ChannelTypeZhipu_v4: "ZhipuV4",
ChannelTypePerplexity: "Perplexity",
ChannelTypeLingYiWanWu: "LingYiWanWu",
ChannelTypeAws: "AWS",
ChannelTypeCohere: "Cohere",
ChannelTypeMiniMax: "MiniMax",
ChannelTypeSunoAPI: "SunoAPI",
ChannelTypeDify: "Dify",
ChannelTypeJina: "Jina",
ChannelCloudflare: "Cloudflare",
ChannelTypeSiliconFlow: "SiliconFlow",
ChannelTypeVertexAi: "VertexAI",
ChannelTypeMistral: "Mistral",
ChannelTypeDeepSeek: "DeepSeek",
ChannelTypeMokaAI: "MokaAI",
ChannelTypeVolcEngine: "VolcEngine",
ChannelTypeBaiduV2: "BaiduV2",
ChannelTypeXinference: "Xinference",
ChannelTypeXai: "xAI",
ChannelTypeCoze: "Coze",
ChannelTypeKling: "Kling",
ChannelTypeJimeng: "Jimeng",
ChannelTypeVidu: "Vidu",
ChannelTypeSubmodel: "Submodel",
ChannelTypeDoubaoVideo: "DoubaoVideo",
ChannelTypeSora: "Sora",
ChannelTypeReplicate: "Replicate",
ChannelTypeCodex: "Codex",
ChannelTypeUnknown: "Unknown",
ChannelTypeOpenAI: "OpenAI",
ChannelTypeMidjourney: "Midjourney",
ChannelTypeAzure: "Azure",
ChannelTypeOllama: "Ollama",
ChannelTypeMidjourneyPlus: "MidjourneyPlus",
ChannelTypeOpenAIMax: "OpenAIMax",
ChannelTypeOhMyGPT: "OhMyGPT",
ChannelTypeCustom: "Custom",
ChannelTypeAILS: "AILS",
ChannelTypeAIProxy: "AIProxy",
ChannelTypePaLM: "PaLM",
ChannelTypeAPI2GPT: "API2GPT",
ChannelTypeAIGC2D: "AIGC2D",
ChannelTypeAnthropic: "Anthropic",
ChannelTypeBaidu: "Baidu",
ChannelTypeZhipu: "Zhipu",
ChannelTypeAli: "Ali",
ChannelTypeXunfei: "Xunfei",
ChannelType360: "360",
ChannelTypeOpenRouter: "OpenRouter",
ChannelTypeAIProxyLibrary: "AIProxyLibrary",
ChannelTypeFastGPT: "FastGPT",
ChannelTypeTencent: "Tencent",
ChannelTypeGemini: "Gemini",
ChannelTypeMoonshot: "Moonshot",
ChannelTypeZhipu_v4: "ZhipuV4",
ChannelTypePerplexity: "Perplexity",
ChannelTypeLingYiWanWu: "LingYiWanWu",
ChannelTypeAws: "AWS",
ChannelTypeCohere: "Cohere",
ChannelTypeMiniMax: "MiniMax",
ChannelTypeSunoAPI: "SunoAPI",
ChannelTypeDify: "Dify",
ChannelTypeJina: "Jina",
ChannelCloudflare: "Cloudflare",
ChannelTypeSiliconFlow: "SiliconFlow",
ChannelTypeVertexAi: "VertexAI",
ChannelTypeMistral: "Mistral",
ChannelTypeDeepSeek: "DeepSeek",
ChannelTypeMokaAI: "MokaAI",
ChannelTypeVolcEngine: "VolcEngine",
ChannelTypeBaiduV2: "BaiduV2",
ChannelTypeXinference: "Xinference",
ChannelTypeXai: "xAI",
ChannelTypeCoze: "Coze",
ChannelTypeKling: "Kling",
ChannelTypeJimeng: "Jimeng",
ChannelTypeVidu: "Vidu",
ChannelTypeSubmodel: "Submodel",
ChannelTypeDoubaoVideo: "DoubaoVideo",
ChannelTypeSora: "Sora",
ChannelTypeReplicate: "Replicate",
ChannelTypeCodex: "Codex",
ChannelTypeDoubaoVideoCompatibleAiping: "DoubaoVideoCompatibleAiping",
ChannelTypeKlingAiping: "KlingAiping",
ChannelTypeDoubaoVideoCompatibleTianyiYun: "DoubaoVideoCompatibleTianyiYun",
}

func GetChannelTypeName(channelType int) string {


+ 2
- 1
constant/context_key.go View File

@@ -15,7 +15,6 @@ const (
ContextKeyTokenKey ContextKey = "token_key"
ContextKeyTokenId ContextKey = "token_id"
ContextKeyTokenGroup ContextKey = "token_group"
ContextKeyTokenSpecificChannelId ContextKey = "specific_channel_id"
ContextKeyTokenModelLimitEnabled ContextKey = "token_model_limit_enabled"
ContextKeyTokenModelLimit ContextKey = "token_model_limit"
ContextKeyTokenCrossGroupRetry ContextKey = "token_cross_group_retry"
@@ -38,6 +37,8 @@ const (
ContextKeyChannelMultiKeyIndex ContextKey = "channel_multi_key_index"
ContextKeyChannelKey ContextKey = "channel_key"

ContextKeyRelayFormat ContextKey = "relay_format"

ContextKeyAutoGroup ContextKey = "auto_group"
ContextKeyAutoGroupIndex ContextKey = "auto_group_index"
ContextKeyAutoGroupRetryIndex ContextKey = "auto_group_retry_index"


+ 21
- 8
constant/endpoint_type.go View File

@@ -1,17 +1,30 @@
package constant

// EndpointType 标识模型支持的 API 端点类型
// 用于模型广场筛选、渠道定价配置、请求路由判断
// 前端筛选组件: web/src/components/table/model-pricing/filter/PricingEndpointTypes.jsx
type EndpointType string

const (
EndpointTypeOpenAI EndpointType = "openai"
EndpointTypeOpenAIResponse EndpointType = "openai-response"
// OpenAI Chat Completions 格式 (/v1/chat/completions)
EndpointTypeOpenAI EndpointType = "openai"
// OpenAI Responses API 格式 (/v1/responses),用于 o3/o4-mini 等仅支持 Responses 的模型
EndpointTypeOpenAIResponse EndpointType = "openai-response"
// OpenAI Responses Compact 格式 (/v1/responses/compact),精简输出
EndpointTypeOpenAIResponseCompact EndpointType = "openai-response-compact"
EndpointTypeAnthropic EndpointType = "anthropic"
EndpointTypeGemini EndpointType = "gemini"
EndpointTypeJinaRerank EndpointType = "jina-rerank"
EndpointTypeImageGeneration EndpointType = "image-generation"
EndpointTypeEmbeddings EndpointType = "embeddings"
EndpointTypeOpenAIVideo EndpointType = "openai-video"
// Anthropic Messages API 格式 (/v1/messages)
EndpointTypeAnthropic EndpointType = "anthropic"
// Google Gemini 原生格式 (/v1beta/models/{model}:generateContent)
EndpointTypeGemini EndpointType = "gemini"
// Jina Rerank API 格式 (/v1/rerank)
EndpointTypeJinaRerank EndpointType = "jina-rerank"
// 图片生成 API (/v1/images/generations),如 DALL-E、Midjourney 代理
EndpointTypeImageGeneration EndpointType = "image-generation"
// 向量嵌入 API (/v1/embeddings)
EndpointTypeEmbeddings EndpointType = "embeddings"
// OpenAI Video API,如 Sora 视频生成
EndpointTypeOpenAIVideo EndpointType = "openai-video"
EndpointTypeDoubaoVideo EndpointType = "doubao-video"
//EndpointTypeMidjourney EndpointType = "midjourney-proxy"
//EndpointTypeSuno EndpointType = "suno-proxy"
//EndpointTypeKling EndpointType = "kling"


+ 2
- 0
controller/channel-test.go View File

@@ -65,6 +65,8 @@ func testChannel(channel *model.Channel, testModel string, endpointType string,
constant.ChannelTypeKling,
constant.ChannelTypeJimeng,
constant.ChannelTypeDoubaoVideo,
constant.ChannelTypeDoubaoVideoCompatibleAiping,
constant.ChannelTypeDoubaoVideoCompatibleTianyiYun,
constant.ChannelTypeVidu,
}
if lo.Contains(unsupportedTestChannelTypes, channel.Type) {


+ 23
- 38
controller/channel.go View File

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

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

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

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

oauthKey, ch, err := service.RefreshCodexChannelCredential(ctx, channelId, service.CodexCredentialRefreshOptions{ResetCaches: true})
if err != nil {
if errors.Is(err, common.ErrCodexOAuthCredentialRequired) {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "当前凭证方式不支持刷新凭证"})
return
}
common.SysError("failed to refresh codex channel credential: " + err.Error())
c.JSON(http.StatusOK, gin.H{"success": false, "message": "刷新凭证失败,请稍后重试"})
return
@@ -2095,29 +2106,3 @@ func OllamaVersion(c *gin.Context) {
})
}

// GetUserChannelsForBinding 获取用户可用于绑定的渠道列表
// 只返回 id、name、type 和 remark,不包含敏感信息(如key)
func GetUserChannelsForBinding(c *gin.Context) {
channels, err := model.GetAllChannelsForBinding()
if err != nil {
common.ApiError(c, err)
return
}

// 构建返回结果
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,
})
}

c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": result,
})
}

+ 0
- 325
controller/channel_pricing.go View File

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

import (
"strconv"
"strings"

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

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

// GetAllChannelPricing 获取所有渠道定价(分页)
func GetAllChannelPricing(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("p", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "10"))
offset := (page - 1) * pageSize

list, total, err := model.GetAllChannelPricing(offset, pageSize)
if err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, gin.H{
"page": page,
"page_size": pageSize,
"total": total,
"items": list,
})
}

// GetChannelPricingByModel 获取指定模型的所有渠道定价
func GetChannelPricingByModel(c *gin.Context) {
modelName := c.Param("name")
if modelName == "" {
common.ApiErrorMsg(c, "model name is required")
return
}

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

common.ApiSuccess(c, list)
}

// 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"`
}

// CreateChannelPricing 创建或更新渠道定价
func CreateChannelPricing(c *gin.Context) {
var req CreateChannelPricingRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
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
if err := existing.Update(); err != nil {
common.ApiError(c, err)
return
}
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,
}
if err := cp.Insert(); err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, cp)
}

// BatchCreateChannelPricingRequest 批量创建请求
type BatchCreateChannelPricingRequest struct {
Items []*CreateChannelPricingRequest `json:"items" binding:"required"`
}

// BatchCreateChannelPricing 批量创建或更新渠道定价
func BatchCreateChannelPricing(c *gin.Context) {
var req BatchCreateChannelPricingRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}

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,
})
}

// 使用事务逐个处理(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
if err := existing.Update(); err != nil {
common.ApiError(c, err)
return
}
} else {
if err := cp.Insert(); err != nil {
common.ApiError(c, err)
return
}
}
}

common.ApiSuccess(c, gin.H{"affected": len(pricings)})
}

// DeleteChannelPricing 删除渠道定价
func DeleteChannelPricing(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.Atoi(idStr)
if err != nil {
common.ApiError(c, err)
return
}

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

common.ApiSuccess(c, nil)
}

// CopyGlobalPricingRequest 复制全局定价请求
type CopyGlobalPricingRequest struct {
Overwrite bool `json:"overwrite"` // 是否覆盖已存在的渠道定价
}

// CopyGlobalPricing 复制全局定价到指定渠道
// 从 ratio_setting 读取全局定价信息,复制到 channel_pricing 表
func CopyGlobalPricing(c *gin.Context) {
channelIdStr := c.Param("channel_id")
channelId, err := strconv.Atoi(channelIdStr)
if err != nil {
common.ApiError(c, err)
return
}

var req CopyGlobalPricingRequest
c.ShouldBindJSON(&req)

// 获取该渠道支持的所有模型(使用 GetAbilitiesByChannelId)
abilities, err := model.GetAbilitiesByChannelId(channelId)
if err != nil {
common.ApiError(c, err)
return
}

imported := 0
for _, ability := range abilities {
// 检查是否已存在
existing, _ := model.GetChannelPricing(ability.Model, channelId)
if existing != nil && !req.Overwrite {
continue
}

// 确定定价类型
var quotaType int
var ratio, completionRatio, price float64

// 优先检查是否有按次计费的价格
modelPrice, hasPrice := ratio_setting.GetModelPrice(ability.Model, false)
if hasPrice {
quotaType = model.QuotaTypeByCall
price = modelPrice
} else {
// 使用按量计费
quotaType = model.QuotaTypeByTokens
modelRatio, hasRatio, _ := ratio_setting.GetModelRatio(ability.Model)
if hasRatio {
ratio = modelRatio
}
completionRatio = ratio_setting.GetCompletionRatio(ability.Model)
}

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

common.ApiSuccess(c, gin.H{
"total": len(abilities),
"imported": imported,
})
}

// ChannelPricingWithTags 带标签详情的渠道定价响应
type ChannelPricingWithTags struct {
*model.ChannelPricing
Tags []*model.PricingTag `json:"tags"`
}

// GetChannelPricingByModelWithChannelInfo 获取指定模型的渠道定价(带渠道信息,所有用户可访问)
func GetChannelPricingByModelWithChannelInfo(c *gin.Context) {
// 使用通配符路由时,参数包含前导斜杠,需要去除
modelName := c.Param("name")
if modelName == "" {
common.ApiErrorMsg(c, "model name is required")
return
}
// 去除前导斜杠(路由是 /channel-pricing/model/*name,name 会是 "/deepseek-ai/xxx")
modelName = strings.TrimPrefix(modelName, "/")

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

common.ApiSuccess(c, list)
}

// GetChannelPricingWithTags 获取渠道定价(带标签详情)
func GetChannelPricingWithTags(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("p", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "10"))
offset := (page - 1) * pageSize

list, total, err := model.GetAllChannelPricing(offset, pageSize)
if err != nil {
common.ApiError(c, err)
return
}

// 获取所有标签
tags, _ := model.GetAllPricingTags()
tagMap := make(map[int]*model.PricingTag)
for _, tag := range tags {
tagMap[tag.Id] = tag
}

// 为每个定价填充标签详情
result := make([]*ChannelPricingWithTags, 0, len(list))
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)
}
}
}
}
result = append(result, item)
}

common.ApiSuccess(c, gin.H{
"page": page,
"page_size": pageSize,
"total": total,
"items": result,
})
}

+ 36
- 0
controller/channel_test_tianyiyun_test.go View File

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

import (
"strings"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
"github.com/stretchr/testify/require"
)

func TestTestChannelRejectsAsyncVideoChannels(t *testing.T) {
channels := []*model.Channel{
{Type: constant.ChannelTypeDoubaoVideoCompatibleAiping, Models: "doubao-seedance-2-0-260128", Status: common.ChannelStatusEnabled},
{Type: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, Models: "cdance2.0-0611", Status: common.ChannelStatusEnabled},
}

for _, channel := range channels {
result := testChannel(channel, "", "", false)

require.Error(t, result.localErr)
require.True(t, strings.Contains(result.localErr.Error(), "channel test is not supported"))
}
}

func TestRequiredTaskChannelTypeForTianyiYunSeedanceModelUsesAllowedFamily(t *testing.T) {
c := newControllerJSONContext(t, "/api/v3/contents/generations/tasks", `{"model":"Doubao-Seedance-2.0"}`)

require.Equal(t, 0, requiredTaskChannelTypeForRequest(c))
require.ElementsMatch(t,
service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilySeedance),
allowedTaskChannelTypesForRequest(c),
)
}

+ 111
- 0
controller/codex_channel_test.go View File

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

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

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

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

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

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

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

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

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

return db
}

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

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

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

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

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

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

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

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

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

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

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

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

+ 8
- 10
controller/codex_usage.go View File

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

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

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

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

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

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

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

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


+ 158
- 0
controller/doubao_aiping_video.go View File

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

import (
"fmt"
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)

func convertAipingNativeVideoRequest(body []byte) (relaycommon.TaskSubmitReq, error) {
var raw map[string]any
if err := common.Unmarshal(body, &raw); err != nil {
return relaycommon.TaskSubmitReq{}, err
}
modelName, _ := raw["model"].(string)
if strings.TrimSpace(modelName) == "" {
return relaycommon.TaskSubmitReq{}, fmt.Errorf("model field is required")
}

metadata := map[string]any{}
for key, value := range raw {
if key == "model" || key == "content" {
continue
}
metadata[key] = value
}

var passthroughContent []any
if rawContent, ok := raw["content"].([]any); ok {
for _, item := range rawContent {
itemMap, ok := item.(map[string]any)
if !ok {
passthroughContent = append(passthroughContent, item)
continue
}
passthroughContent = append(passthroughContent, itemMap)
}
}
if len(passthroughContent) > 0 {
metadata["content"] = passthroughContent
}

return relaycommon.TaskSubmitReq{
Model: strings.TrimSpace(modelName),
Metadata: metadata,
}, nil
}

func prepareAipingNativeVideoSubmit(c *gin.Context, info *relaycommon.RelayInfo) error {
bodyStorage, err := common.GetBodyStorage(c)
if err != nil {
return err
}
body, err := bodyStorage.Bytes()
if err != nil {
return err
}
req, err := convertAipingNativeVideoRequest(body)
if err != nil {
return err
}
info.OriginModelName = req.Model
info.Action = constant.TaskActionGenerate
relaycommon.StoreTaskRequest(c, info, constant.TaskActionGenerate, req)
return nil
}

func AipingNativeVideoSubmit(c *gin.Context) {
relayInfo, err := relaycommon.GenRelayInfo(c, types.RelayFormatTask, nil, nil)
if err != nil {
c.JSON(http.StatusInternalServerError, &dto.TaskError{
Code: "gen_relay_info_failed",
Message: err.Error(),
StatusCode: http.StatusInternalServerError,
})
return
}
if relayInfo.ChannelMeta == nil {
relayInfo.ChannelMeta = &relaycommon.ChannelMeta{}
}
if err := prepareAipingNativeVideoSubmit(c, relayInfo); err != nil {
c.JSON(http.StatusBadRequest, &dto.TaskError{
Code: "invalid_request",
Message: err.Error(),
StatusCode: http.StatusBadRequest,
})
return
}
relayTaskWithInfo(c, relayInfo)
}

func buildAipingNativeFetchResponse(task *model.Task) ([]byte, error) {
payload := map[string]any{}
if len(task.Data) > 0 {
if err := common.Unmarshal(task.Data, &payload); err != nil {
return nil, err
}
}
delete(payload, "aiping_id")
if _, ok := payload["error"]; !ok {
payload["error"] = nil
}
payload["id"] = task.TaskID
if _, ok := payload["model"]; !ok {
payload["model"] = task.Properties.OriginModelName
}
if _, ok := payload["status"]; !ok {
payload["status"] = mapAipingNativeStatus(task.Status)
}
if _, ok := payload["created_at"]; !ok {
payload["created_at"] = task.CreatedAt
}
if _, ok := payload["updated_at"]; !ok {
payload["updated_at"] = task.UpdatedAt
}
return common.Marshal(payload)
}

func mapAipingNativeStatus(status model.TaskStatus) string {
switch status {
case model.TaskStatusQueued, model.TaskStatusSubmitted:
return "queued"
case model.TaskStatusInProgress:
return "running"
case model.TaskStatusSuccess:
return "succeeded"
case model.TaskStatusFailure:
return "failed"
default:
return "running"
}
}

func AipingNativeVideoFetch(c *gin.Context) {
taskID := c.Param("task_id")
task, exist, err := model.GetByTaskId(c.GetInt("id"), taskID)
if err != nil {
c.JSON(http.StatusInternalServerError, &dto.TaskError{Code: "get_task_failed", Message: err.Error(), StatusCode: http.StatusInternalServerError})
return
}
if !exist {
c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"code": "not_found", "message": "task not found"}})
return
}
data, err := buildAipingNativeFetchResponse(task)
if err != nil {
c.JSON(http.StatusInternalServerError, &dto.TaskError{Code: "build_response_failed", Message: err.Error(), StatusCode: http.StatusInternalServerError})
return
}
c.Data(http.StatusOK, "application/json", data)
}

+ 209
- 0
controller/doubao_aiping_video_test.go View File

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

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

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestConvertAipingNativeVideoRequestPreservesNativeContentOrder(t *testing.T) {
req, err := convertAipingNativeVideoRequest([]byte(`{
"model":"doubao-seedance-2-0-260128",
"content":[
{"type":"text","text":"first prompt"},
{"type":"image_url","image_url":{"url":"asset://img1"},"role":"reference_image"},
{"type":"text","text":"second prompt"},
{"type":"video_url","video_url":{"url":"https://example.test/input.mp4"},"role":"reference_video"}
],
"duration":5,
"resolution":"480p",
"tools":[{"type":"web_search"}]
}`))

require.NoError(t, err)
require.Equal(t, "doubao-seedance-2-0-260128", req.Model)
require.Empty(t, req.Prompt)
require.Equal(t, float64(5), req.Metadata["duration"])
require.Equal(t, "480p", req.Metadata["resolution"])
require.NotNil(t, req.Metadata["tools"])
content, ok := req.Metadata["content"].([]any)
require.True(t, ok)
require.Len(t, content, 4)
require.Equal(t, "text", content[0].(map[string]any)["type"])
require.Equal(t, "image_url", content[1].(map[string]any)["type"])
require.Equal(t, "text", content[2].(map[string]any)["type"])
require.Equal(t, "video_url", content[3].(map[string]any)["type"])
}

func TestConvertAipingNativeVideoRequestPreservesArkDocumentFields(t *testing.T) {
req, err := convertAipingNativeVideoRequest([]byte(`{
"model":"doubao-seedance-2-0-260128",
"content":[
{"type":"text","text":"开场:海边日落"},
{"type":"image_url","image_url":{"url":"asset://image-1"},"role":"reference_image"},
{"type":"video_url","video_url":{"url":"asset://video-1"},"role":"reference_video"},
{"type":"audio_url","audio_url":{"url":"asset://audio-1"},"role":"reference_audio"},
{"type":"draft_task","draft_task":{"id":"cgt-draft"}}
],
"callback_url":"https://example.test/callback",
"return_last_frame":true,
"service_tier":"default",
"execution_expires_after":3600,
"generate_audio":false,
"draft":true,
"tools":[{"type":"web_search"}],
"safety_identifier":"user-hash-1",
"priority":5,
"resolution":"480p",
"ratio":"1:1",
"duration":5,
"frames":29,
"seed":11,
"camera_fixed":false,
"watermark":true
}`))

require.NoError(t, err)
require.Equal(t, "doubao-seedance-2-0-260128", req.Model)
require.Empty(t, req.Prompt)
require.Equal(t, "https://example.test/callback", req.Metadata["callback_url"])
require.Equal(t, true, req.Metadata["return_last_frame"])
require.Equal(t, "default", req.Metadata["service_tier"])
require.Equal(t, float64(3600), req.Metadata["execution_expires_after"])
require.Equal(t, false, req.Metadata["generate_audio"])
require.Equal(t, true, req.Metadata["draft"])
require.Equal(t, "user-hash-1", req.Metadata["safety_identifier"])
require.Equal(t, float64(5), req.Metadata["priority"])
require.Equal(t, "480p", req.Metadata["resolution"])
require.Equal(t, "1:1", req.Metadata["ratio"])
require.Equal(t, float64(5), req.Metadata["duration"])
require.Equal(t, float64(29), req.Metadata["frames"])
require.Equal(t, float64(11), req.Metadata["seed"])
require.Equal(t, false, req.Metadata["camera_fixed"])
require.Equal(t, true, req.Metadata["watermark"])

content, ok := req.Metadata["content"].([]any)
require.True(t, ok)
require.Len(t, content, 5)
require.Equal(t, "audio_url", content[3].(map[string]any)["type"])
require.Equal(t, "draft_task", content[4].(map[string]any)["type"])
tools, ok := req.Metadata["tools"].([]any)
require.True(t, ok)
require.Equal(t, "web_search", tools[0].(map[string]any)["type"])
}

func TestConvertAipingNativeVideoRequestRejectsMissingModel(t *testing.T) {
_, err := convertAipingNativeVideoRequest([]byte(`{"content":[{"type":"text","text":"prompt"}]}`))
require.ErrorContains(t, err, "model")
}

func TestPrepareAipingNativeVideoSubmitStoresConvertedRequest(t *testing.T) {
c := newControllerJSONContext(t, "/api/v3/contents/generations/tasks", `{
"model":"doubao-seedance-2-0-260128",
"content":[{"type":"text","text":"prompt"}]
}`)
info := &relaycommon.RelayInfo{TaskRelayInfo: &relaycommon.TaskRelayInfo{}}

err := prepareAipingNativeVideoSubmit(c, info)

require.NoError(t, err)
stored, err := relaycommon.GetTaskRequest(c)
require.NoError(t, err)
require.Equal(t, "doubao-seedance-2-0-260128", stored.Model)
require.Empty(t, stored.Prompt)
content, ok := stored.Metadata["content"].([]any)
require.True(t, ok)
require.Equal(t, "prompt", content[0].(map[string]any)["text"])
require.Equal(t, constant.TaskActionGenerate, info.Action)
require.Equal(t, "doubao-seedance-2-0-260128", info.OriginModelName)
}

func TestBuildAipingNativeFetchResponseRewritesIDAndPreservesUpstreamFields(t *testing.T) {
task := &model.Task{
TaskID: "task_public",
Status: model.TaskStatusSuccess,
CreatedAt: 1781496040,
UpdatedAt: 1781496278,
Properties: model.Properties{OriginModelName: "doubao-seedance-2-0-260128"},
Data: []byte(`{
"id":"cgt-upstream",
"aiping_id":"3d2c8c17-36ad-4138-8602-88b3c60e56c6",
"model":"doubao-seedance-2-0-260128",
"status":"succeeded",
"content":{"video_url":"https://example.test/output.mp4"},
"usage":{"completion_tokens":48400,"total_tokens":48400},
"created_at":1781496040,
"updated_at":1781496278,
"seed":73812,
"resolution":"480p",
"ratio":"1:1",
"duration":5,
"framespersecond":24,
"service_tier":"default",
"execution_expires_after":172800,
"generate_audio":true,
"draft":false,
"priority":0
}`),
}

data, err := buildAipingNativeFetchResponse(task)

require.NoError(t, err)
var payload map[string]any
require.NoError(t, common.Unmarshal(data, &payload))
require.Equal(t, "task_public", payload["id"])
require.Equal(t, "succeeded", payload["status"])
require.Equal(t, "doubao-seedance-2-0-260128", payload["model"])
require.Equal(t, "https://example.test/output.mp4", payload["content"].(map[string]any)["video_url"])
require.Equal(t, float64(48400), payload["usage"].(map[string]any)["total_tokens"])
require.Equal(t, float64(73812), payload["seed"])
require.Equal(t, "480p", payload["resolution"])
require.Equal(t, "1:1", payload["ratio"])
require.Equal(t, float64(5), payload["duration"])
require.Equal(t, float64(24), payload["framespersecond"])
require.Equal(t, "default", payload["service_tier"])
require.Equal(t, float64(172800), payload["execution_expires_after"])
require.Equal(t, true, payload["generate_audio"])
require.Equal(t, false, payload["draft"])
require.Equal(t, float64(0), payload["priority"])
require.Contains(t, payload, "error")
require.Nil(t, payload["error"])
require.NotContains(t, payload, "aiping_id")
}

func TestBuildAipingNativeFetchResponseFallsBackForFailedTask(t *testing.T) {
task := &model.Task{
TaskID: "task_public",
Status: model.TaskStatusFailure,
CreatedAt: 100,
UpdatedAt: 200,
Properties: model.Properties{OriginModelName: "doubao-seedance-2-0-260128"},
Data: []byte(`{"id":"cgt-upstream","status":"failed","error":{"code":"InvalidParameter","message":"duration is invalid"}}`),
}

data, err := buildAipingNativeFetchResponse(task)

require.NoError(t, err)
require.Contains(t, string(data), `"id":"task_public"`)
require.Contains(t, string(data), `"code":"InvalidParameter"`)
require.Contains(t, string(data), `"message":"duration is invalid"`)
require.NotContains(t, string(data), `"error":null`)
}

func newControllerJSONContext(t *testing.T, path string, body string) *gin.Context {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
return c
}

+ 203
- 0
controller/doubao_asset.go View File

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

import (
"context"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/system_setting"

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

const defaultDoubaoAssetBaseURL = "https://ark.cn-beijing.volcengineapi.com"

var blockedAssetActions = map[string]struct{}{
"createassetgroup": {},
"getassetgroup": {},
"listassetgroups": {},
"updateassetgroup": {},
"deleteassetgroup": {},
}

func assetProxyError(c *gin.Context, status int, errType, message string) {
c.JSON(status, gin.H{
"error": gin.H{
"message": message,
"type": errType,
},
})
}

func buildDoubaoAssetURL(channel *model.Channel, action string, version string) (string, error) {
baseURL := defaultDoubaoAssetBaseURL
if channel != nil && channel.BaseURL != nil {
if configuredBaseURL := strings.TrimSpace(*channel.BaseURL); configuredBaseURL != "" {
baseURL = configuredBaseURL
}
}

u, err := url.Parse(baseURL)
if err != nil {
return "", err
}
if u.Scheme == "" || u.Host == "" {
return "", fmt.Errorf("invalid Doubao asset base URL: %s", baseURL)
}

u.Path = strings.TrimRight(u.Path, "/") + "/api/v1/multimodal/sd/assets"
u.RawQuery = ""
u.Fragment = ""

query := u.Query()
query.Set("Action", action)
if strings.TrimSpace(version) == "" {
version = "2024-01-01"
}
query.Set("Version", version)
u.RawQuery = query.Encode()

return u.String(), nil
}

func effectiveDoubaoAssetGroup(c *gin.Context) string {
if group := strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyUsingGroup)); group != "" {
return group
}
return strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyTokenGroup))
}

func concreteDoubaoAssetGroupsForRequest(c *gin.Context, autoGroups func(string) []string) []string {
group := effectiveDoubaoAssetGroup(c)
if group != "auto" {
if group == "" {
return nil
}
return []string{group}
}

userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
groups := autoGroups(userGroup)
concreteGroups := make([]string, 0, len(groups))
for _, candidate := range groups {
candidate = strings.TrimSpace(candidate)
if candidate == "" || candidate == "auto" {
continue
}
concreteGroups = append(concreteGroups, candidate)
}
return concreteGroups
}

func resolveDoubaoAssetChannelForRequest(c *gin.Context) (*model.Channel, string, error) {
userId := c.GetInt("id")
var lastErr error
for _, group := range concreteDoubaoAssetGroupsForRequest(c, service.GetUserAutoGroup) {
channel, err := service.ResolveDoubaoAssetChannel(userId, group)
if err != nil {
lastErr = err
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to resolve Doubao asset channel for group %s: %v", group, err))
continue
}
if channel != nil {
return channel, group, nil
}
}
if lastErr != nil {
return nil, "", lastErr
}
return nil, "", nil
}

func DoubaoAssetProxy(c *gin.Context) {
action := strings.TrimSpace(c.Query("Action"))
if action == "" {
assetProxyError(c, http.StatusBadRequest, "invalid_request_error", "Action query parameter is required")
return
}
if _, ok := blockedAssetActions[strings.ToLower(action)]; ok {
assetProxyError(c, http.StatusBadRequest, "invalid_request_error", fmt.Sprintf("Asset group API (%s) is not supported", action))
return
}

version := strings.TrimSpace(c.Query("Version"))
if version == "" {
version = "2024-01-01"
}

channel, _, err := resolveDoubaoAssetChannelForRequest(c)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to resolve Doubao asset channel: %v", err))
return
}
if channel == nil {
assetProxyError(c, http.StatusBadGateway, "server_error", "Failed to resolve Doubao asset channel")
return
}
if strings.TrimSpace(channel.Key) == "" {
assetProxyError(c, http.StatusBadGateway, "server_error", "Doubao asset channel key is missing")
return
}

upstreamURL, err := buildDoubaoAssetURL(channel, action, version)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to build upstream URL: %v", err))
return
}

fetchSetting := system_setting.GetFetchSetting()
if err := common.ValidateURLWithFetchSetting(upstreamURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Doubao asset URL blocked: %v", err))
assetProxyError(c, http.StatusForbidden, "server_error", fmt.Sprintf("request blocked: %v", err))
return
}

client, err := service.GetHttpClientWithProxy(channel.GetSetting().Proxy)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to create proxy client: %v", err))
return
}
if client == nil {
assetProxyError(c, http.StatusBadGateway, "server_error", "Failed to create proxy client")
return
}

ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
defer cancel()

req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, c.Request.Body)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to create upstream request: %v", err))
return
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("Authorization", "Bearer "+strings.TrimSpace(channel.Key))

resp, err := client.Do(req)
if err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to proxy Doubao asset request to %s: %s", upstreamURL, err.Error()))
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to proxy Doubao asset request: %v", err))
return
}
defer resp.Body.Close()

for key, values := range resp.Header {
for _, value := range values {
c.Writer.Header().Add(key, value)
}
}
c.Writer.WriteHeader(resp.StatusCode)
if _, err = io.Copy(c.Writer, resp.Body); err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to copy Doubao asset upstream response: %s", err.Error()))
}
}

+ 103
- 0
controller/doubao_asset_test.go View File

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

import (
"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/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func setupDoubaoAssetProxyRouter(t *testing.T) *gin.Engine {
t.Helper()

oldMode := gin.Mode()
gin.SetMode(gin.TestMode)
t.Cleanup(func() {
gin.SetMode(oldMode)
})

r := gin.New()
r.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)
return r
}

func decodeDoubaoAssetErrorMessage(t *testing.T, body string) string {
t.Helper()

var payload struct {
Error struct {
Message string `json:"message"`
Type string `json:"type"`
} `json:"error"`
}
require.NoError(t, common.Unmarshal([]byte(body), &payload))
return payload.Error.Message
}

func TestDoubaoAssetProxyMissingActionReturns400(t *testing.T) {
router := setupDoubaoAssetProxyRouter(t)

req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Version=2024-01-01", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Equal(t, "Action query parameter is required", decodeDoubaoAssetErrorMessage(t, w.Body.String()))
}

func TestDoubaoAssetProxyBlocksAssetGroupActionsCaseInsensitively(t *testing.T) {
router := setupDoubaoAssetProxyRouter(t)

for _, action := range []string{"CreateAssetGroup", "getassetgroup", "LISTASSETGROUPS", "UpdateAssetGroup", "deleteassetgroup"} {
t.Run(action, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+action, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, decodeDoubaoAssetErrorMessage(t, w.Body.String()), "Asset group API ("+action+") is not supported")
})
}
}

func TestBuildDoubaoAssetURLDefaultBaseAndEscapedQuery(t *testing.T) {
got, err := buildDoubaoAssetURL(&model.Channel{}, "ApplyUploadInner&Space", "")

require.NoError(t, err)
assert.Equal(t, defaultDoubaoAssetBaseURL+"/api/v1/multimodal/sd/assets?Action=ApplyUploadInner%26Space&Version=2024-01-01", got)
}

func TestBuildDoubaoAssetURLExplicitBaseURLOverridesDefault(t *testing.T) {
baseURL := "https://example.com/custom/"

got, err := buildDoubaoAssetURL(&model.Channel{BaseURL: &baseURL}, "CommitUploadInner", "2025-02-03")

require.NoError(t, err)
assert.Equal(t, "https://example.com/custom/api/v1/multimodal/sd/assets?Action=CommitUploadInner&Version=2025-02-03", got)
}

func TestEffectiveDoubaoAssetGroupUsesUsingGroupBeforeBlankTokenGroup(t *testing.T) {
c, _ := gin.CreateTestContext(httptest.NewRecorder())
common.SetContextKey(c, constant.ContextKeyTokenGroup, "")
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")

assert.Equal(t, "default", effectiveDoubaoAssetGroup(c))
}

func TestConcreteDoubaoAssetGroupsForAutoUsesUserAutoGroups(t *testing.T) {
c, _ := gin.CreateTestContext(httptest.NewRecorder())
common.SetContextKey(c, constant.ContextKeyUsingGroup, "auto")
common.SetContextKey(c, constant.ContextKeyTokenGroup, "")

groups := concreteDoubaoAssetGroupsForRequest(c, func(string) []string {
return []string{"default", "vip"}
})

assert.Equal(t, []string{"default", "vip"}, groups)
}

+ 128
- 0
controller/email_quota_rule.go View File

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

import (
"strconv"

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

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

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

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

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

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

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

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

common.ApiSuccess(c, rule)
}

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

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

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

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

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

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

common.ApiSuccess(c, rule)
}

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

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

common.ApiSuccess(c, nil)
}

+ 264
- 0
controller/email_quota_rule_test.go View File

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

+ 203
- 0
controller/internal_user_migration.go View File

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

import (
"errors"
"net/http"
"strconv"

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

func ListMigrationUsers(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "0"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
keyword := c.Query("keyword")
if pageSize <= 0 || pageSize > 200 {
pageSize = 50
}

users, total, err := model.QueryLocalUsersForMigration(page, pageSize, keyword)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "database error"})
return
}

snapshots := make([]user_migration.RemoteUserSnapshot, 0, len(users))
for _, user := range users {
bindings, _ := model.GetUserOAuthBindingsByUserId(user.Id)
snapshots = append(snapshots, buildRemoteUserSnapshot(user, bindings))
}
c.JSON(http.StatusOK, gin.H{"success": true, "data": snapshots, "total": total})
}

func QueryMigrationUsers(c *gin.Context) {
var req user_migration.QueryRemoteUsersRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid request body"})
return
}
if req.SelectionMode != user_migration.SelectionModeExplicitIDs {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid selection_mode"})
return
}
if req.PageSize <= 0 || req.PageSize > 200 {
req.PageSize = 50
}

indexed := make(map[int]struct{}, len(req.SourceUserIDs))
orderedIDs := make([]int, 0, len(req.SourceUserIDs))
for _, id := range req.SourceUserIDs {
if id <= 0 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid source_user_ids"})
return
}
if _, ok := indexed[id]; ok {
continue
}
indexed[id] = struct{}{}
orderedIDs = append(orderedIDs, id)
}

summary := &user_migration.SelectionSummary{
Requested: len(orderedIDs),
Excluded: make([]user_migration.SelectionExcluded, 0),
}
snapshots := make([]user_migration.RemoteUserSnapshot, 0, len(orderedIDs))

for _, userID := range orderedIDs {
user, err := model.GetUserById(userID, true)
if err != nil || user == nil {
summary.Excluded = append(summary.Excluded, user_migration.SelectionExcluded{UserID: userID, Reason: "not_found"})
continue
}
if user.Role == common.RoleRootUser {
summary.Excluded = append(summary.Excluded, user_migration.SelectionExcluded{UserID: userID, Reason: "root_user_not_migratable"})
continue
}
if user.Source == common.UserSourceSynced {
summary.Excluded = append(summary.Excluded, user_migration.SelectionExcluded{UserID: userID, Reason: "already_synced"})
continue
}

bindings, _ := model.GetUserOAuthBindingsByUserId(userID)
snapshots = append(snapshots, buildRemoteUserSnapshot(user, bindings))
}

summary.Matched = len(snapshots)
start := req.Page * req.PageSize
if start > len(snapshots) {
start = len(snapshots)
}
end := start + req.PageSize
if end > len(snapshots) {
end = len(snapshots)
}

c.JSON(http.StatusOK, gin.H{
"success": true,
"data": snapshots[start:end],
"total": len(snapshots),
"selection_summary": summary,
})
}

func GetMigrationUser(c *gin.Context) {
userId, err := strconv.Atoi(c.Param("id"))
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid user id"})
return
}

user, err := model.GetUserById(userId, true)
if err != nil || user == nil {
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "user not found"})
return
}
if user.Role == common.RoleRootUser {
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "user not found"})
return
}

bindings, _ := model.GetUserOAuthBindingsByUserId(userId)
c.JSON(http.StatusOK, gin.H{"success": true, "data": buildRemoteUserSnapshot(user, bindings)})
}

func ConvertMigrationUserToSynced(c *gin.Context) {
userId, err := strconv.Atoi(c.Param("id"))
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid user id"})
return
}

var req user_migration.ConvertRemoteUserRequest
if err := c.ShouldBindJSON(&req); err != nil || req.RemoteUserId == 0 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid request body"})
return
}

if err := model.ConvertUserToSynced(userId, req.RemoteUserId, req.SyncedQuota); err != nil {
if errors.Is(err, model.ErrConvertRootUserToSynced) ||
errors.Is(err, model.ErrSyncedUserRemoteUserIDImmutable) {
c.JSON(http.StatusConflict, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true})
}

func CheckSyncedCopy(c *gin.Context) {
cnUserId, err := strconv.Atoi(c.Query("cn_user_id"))
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid cn_user_id"})
return
}

hasCopy, err := model.HasOVSyncedCopyByRemoteId(cnUserId)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "has_copy": hasCopy})
}

func buildRemoteUserSnapshot(user *model.User, bindings []*model.UserOAuthBinding) user_migration.RemoteUserSnapshot {
snapshot := user_migration.RemoteUserSnapshot{
Id: user.Id,
Username: user.Username,
Password: user.Password,
Email: user.Email,
DisplayName: user.DisplayName,
Status: user.Status,
Role: user.Role,
Group: user.Group,
Quota: user.Quota,
AffCode: user.AffCode,
CreatedAt: user.CreatedAt,
Setting: user.Setting,
GitHubId: user.GitHubId,
DiscordId: user.DiscordId,
OidcId: user.OidcId,
WeChatId: user.WeChatId,
TelegramId: user.TelegramId,
LinuxDOId: user.LinuxDOId,
Source: user.Source,
RemoteUserId: user.RemoteUserId,
SyncedQuota: user.SyncedQuota,
}
for _, binding := range bindings {
provider, err := model.GetCustomOAuthProviderById(binding.ProviderId)
if err != nil {
continue
}
snapshot.OAuthBindings = append(snapshot.OAuthBindings, user_migration.RemoteOAuthBinding{
ProviderSlug: provider.Slug,
ProviderUserId: binding.ProviderUserId,
})
}
return snapshot
}

+ 491
- 0
controller/internal_user_migration_test.go View File

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

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

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

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

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

sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

require.NoError(t, db.AutoMigrate(&model.User{}, &model.CustomOAuthProvider{}, &model.UserOAuthBinding{}))

return db
}

func TestListMigrationUsers_ReturnsLocalOnly(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users", ListMigrationUsers)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000001,
Username: "ov-1",
Password: "hash",
Email: "u1@example.com",
AffCode: "A1",
Source: common.UserSourceLocal,
Quota: 99,
}).Error)
require.NoError(t, db.Create(&model.User{
Id: 10000002,
Username: "ov-2",
Password: "hash",
AffCode: "A2",
Source: common.UserSourceSynced,
RemoteUserId: 1,
}).Error)
require.NoError(t, db.Create(&model.User{
Id: 10000003,
Username: "root",
Password: "hash",
AffCode: "A3",
Source: common.UserSourceLocal,
Role: common.RoleRootUser,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/internal/migration/users?page=0&page_size=20", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var body map[string]any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
require.True(t, body["success"].(bool))
require.EqualValues(t, 1, body["total"])
}

func TestListMigrationUsers_FiltersByKeyword(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users", ListMigrationUsers)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000011,
Username: "alice-local",
Password: "hash",
Email: "alice@example.com",
AffCode: "A11",
Source: common.UserSourceLocal,
}).Error)
require.NoError(t, db.Create(&model.User{
Id: 10000012,
Username: "bob-local",
Password: "hash",
Email: "bob@example.com",
AffCode: "A12",
Source: common.UserSourceLocal,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/internal/migration/users?page=0&page_size=20&keyword=alice", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var body struct {
Success bool `json:"success"`
Data []user_migration.RemoteUserSnapshot `json:"data"`
Total int `json:"total"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
require.True(t, body.Success)
require.Equal(t, 1, body.Total)
require.Len(t, body.Data, 1)
require.Equal(t, 10000011, body.Data[0].Id)
}

func TestQueryMigrationUsers_OnlyReturnsSelectedLocalUsers(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/query", QueryMigrationUsers)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{Id: 10000001, Username: "u1", Password: "hash", AffCode: "A1", Source: common.UserSourceLocal}).Error)
require.NoError(t, db.Create(&model.User{Id: 10000002, Username: "u2", Password: "hash", AffCode: "A2", Source: common.UserSourceLocal, Role: common.RoleRootUser}).Error)
require.NoError(t, db.Create(&model.User{Id: 10000003, Username: "u3", Password: "hash", AffCode: "A3", Source: common.UserSourceSynced}).Error)

req := httptest.NewRequest(http.MethodPost, "/api/internal/migration/users/query", strings.NewReader(`{"page":0,"page_size":50,"selection_mode":"explicit_ids","source_user_ids":[10000001,10000002,10000003,10000004]}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var body struct {
Success bool `json:"success"`
Data []user_migration.RemoteUserSnapshot `json:"data"`
Total int `json:"total"`
SelectionSummary user_migration.SelectionSummary `json:"selection_summary"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
require.True(t, body.Success)
require.Len(t, body.Data, 1)
require.Equal(t, 10000001, body.Data[0].Id)
require.Equal(t, 1, body.Total)
require.Equal(t, 4, body.SelectionSummary.Requested)
require.Equal(t, 1, body.SelectionSummary.Matched)
require.Len(t, body.SelectionSummary.Excluded, 3)
}

func TestQueryMigrationUsers_InvalidSelectionModeRejected(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/query", QueryMigrationUsers)

req := httptest.NewRequest(http.MethodPost, "/api/internal/migration/users/query", strings.NewReader(`{"page":0,"page_size":50,"selection_mode":"all","source_user_ids":[1]}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusBadRequest, w.Code)
}

func TestGetMigrationUser_RootExcluded(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users/:id", GetMigrationUser)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000011,
Username: "root",
Password: "hash",
AffCode: "ROOT1",
Source: common.UserSourceLocal,
Role: common.RoleRootUser,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/internal/migration/users/10000011", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusNotFound, w.Code)
}

func TestGetMigrationUser_SyncedUserStillReadableForVerify(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users/:id", GetMigrationUser)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000021,
Username: "synced-user",
Password: "hash",
AffCode: "SYNC1",
Source: common.UserSourceSynced,
RemoteUserId: 88,
SyncedQuota: 123,
Role: common.RoleCommonUser,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/internal/migration/users/10000021", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

func TestGetMigrationUser_InvalidID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users/:id", GetMigrationUser)

req := httptest.NewRequest(http.MethodGet, "/api/internal/migration/users/not-a-number", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusBadRequest, w.Code)
}

func TestConvertMigrationUserToSynced_Idempotent(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000002,
Username: "ov-2",
Password: "hash",
AffCode: "A2",
Source: common.UserSourceLocal,
Quota: 66,
}).Error)

body := `{"remote_user_id":2002,"synced_quota":66}`
for i := 0; i < 2; i++ {
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000002/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
}
}

func TestConvertMigrationUserToSynced_ConflictRemoteId(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000003,
Username: "ov-3",
Password: "hash",
AffCode: "A3",
Source: common.UserSourceSynced,
RemoteUserId: 5000,
}).Error)

body := `{"remote_user_id":9999,"synced_quota":100}`
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000003/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)
}

func TestConvertMigrationUserToSynced_ConflictUsesSentinelErrors(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000013,
Username: "ov-13",
Password: "hash",
AffCode: "A13",
Source: common.UserSourceSynced,
RemoteUserId: 5001,
}).Error)

body := `{"remote_user_id":9999,"synced_quota":100}`
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000013/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)
}

func TestConvertMigrationUserToSynced_RejectsRootUser(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000012,
Username: "root",
Password: "hash",
AffCode: "ROOT2",
Source: common.UserSourceLocal,
Role: common.RoleRootUser,
}).Error)

body := `{"remote_user_id":9999,"synced_quota":100}`
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000012/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)
}

func TestConvertMigrationUserToSynced_InvalidRemoteUserID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000014,
Username: "ov-14",
Password: "hash",
AffCode: "A14",
Source: common.UserSourceLocal,
}).Error)

body := `{"remote_user_id":0,"synced_quota":100}`
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000014/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusBadRequest, w.Code)
}

func TestConvertMigrationUserToSynced_PersistsSyncedFields(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.POST("/api/internal/migration/users/:id/convert-to-synced", ConvertMigrationUserToSynced)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000015,
Username: "ov-15",
Password: "hash",
AffCode: "A15",
Source: common.UserSourceLocal,
Quota: 88,
}).Error)

body := `{"remote_user_id":2015,"synced_quota":166}`
req := httptest.NewRequest(http.MethodPost,
"/api/internal/migration/users/10000015/convert-to-synced",
strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)

var saved model.User
require.NoError(t, db.First(&saved, 10000015).Error)
require.Equal(t, common.UserSourceSynced, saved.Source)
require.Equal(t, 2015, saved.RemoteUserId)
require.Equal(t, 166, saved.SyncedQuota)
require.Greater(t, saved.LastSyncAt, int64(0))
}

func TestCheckSyncedCopy(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testSyncAuthMiddleware())
router.GET("/api/internal/migration/users/check-synced-copy", CheckSyncedCopy)

db := setupInternalMigrationDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.User{
Id: 10000010,
Username: "synced",
Password: "hash",
AffCode: "S1",
Source: common.UserSourceSynced,
RemoteUserId: 456,
}).Error)

req := httptest.NewRequest(http.MethodGet,
"/api/internal/migration/users/check-synced-copy?cn_user_id=456", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)

var body map[string]any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
require.True(t, body["has_copy"].(bool))
}

+ 510
- 0
controller/kling_aiping_native.go View File

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

import (
"bytes"
"fmt"
"io"
"net/http"
"strconv"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/middleware"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/relay"
klingaiping "github.com/QuantumNous/new-api/relay/channel/task/kling/aiping"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)

func KlingAipingNativeTaskSubmit(c *gin.Context) {
route, ok := klingaiping.FindRoute(c.Request.Method, c.FullPath(), klingaiping.RouteKindSubmit)
if !ok {
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "route not found"})
return
}

relayInfo, err := relaycommon.GenRelayInfo(c, types.RelayFormatTask, nil, nil)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": err.Error()})
return
}
payload, err := readJSONPayload(c)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": err.Error()})
return
}

configureKlingAipingTaskRelayInfo(c, relayInfo, route, payload)

if bound := resolveKlingAipingBoundChannel(c, relayInfo); bound != nil {
relayInfo.LockedChannel = bound
}

relayTaskWithInfo(c, relayInfo)
}

func resolveKlingAipingBoundChannel(c *gin.Context, relayInfo *relaycommon.RelayInfo) *model.Channel {
group := concreteTaskVideoBindingGroup(c, relayInfo.TokenGroup)
if group == "" {
return nil
}
channel, err := service.GetBoundKlingAssetChannelForModel(c.GetInt("id"), group, relayInfo.OriginModelName)
if err != nil {
logger.LogError(c, fmt.Sprintf("resolve kling aiping bound channel failed: %v", err))
return nil
}
if channel == nil || !service.IsUsableKlingAssetChannel(channel, group) {
return nil
}
if !model.IsChannelEnabledForGroupModel(group, relayInfo.OriginModelName, channel.Id) {
return nil
}
return channel
}

func configureKlingAipingTaskRelayInfo(c *gin.Context, relayInfo *relaycommon.RelayInfo, route klingaiping.Route, payload map[string]any) {
// Native Kling routes bypass Distribute(), so force relayTaskWithInfo to
// select a channel instead of reading a preselected one from context.
if relayInfo.ChannelMeta == nil {
relayInfo.ChannelMeta = &relaycommon.ChannelMeta{}
}
modelName := resolveKlingAipingModel(payload, route)
relayInfo.OriginModelName = modelName
relayInfo.Action = route.Action
relaycommon.StoreTaskRequest(c, relayInfo, route.Action, relaycommon.TaskSubmitReq{
Model: modelName,
Prompt: stringFromMap(payload, "prompt"),
Duration: durationFromMap(payload),
Metadata: payload,
})
}

func durationFromMap(m map[string]any) int {
for _, key := range []string{"duration", "seconds"} {
if duration, ok := intFromMapValue(m[key]); ok {
return duration
}
}
return 0
}

func intFromMapValue(value any) (int, bool) {
switch v := value.(type) {
case int:
return v, true
case int64:
return int(v), true
case float64:
return int(v), true
case string:
duration, err := strconv.Atoi(strings.TrimSpace(v))
if err == nil {
return duration, true
}
}
return 0, false
}

func KlingAipingNativeTaskFetch(c *gin.Context) {
route, ok := klingaiping.FindRoute(c.Request.Method, c.FullPath(), klingaiping.RouteKindFetch)
if !ok {
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "route not found"})
return
}
task, ok := getKlingAipingUserTask(c, c.Param("task_id"), route.Action)
if !ok {
return
}
refreshKlingAipingTaskIfNeeded(task)
c.JSON(http.StatusOK, buildKlingAipingTaskPayload(task))
}

func KlingAipingNativeTaskList(c *gin.Context) {
route, ok := klingaiping.FindRoute(c.Request.Method, c.FullPath(), klingaiping.RouteKindList)
if !ok {
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "route not found"})
return
}
pageNum, pageSize, err := parseKlingAipingPage(c)
if err != nil {
c.JSON(http.StatusUnprocessableEntity, gin.H{"code": 422, "message": err.Error()})
return
}
tasks := model.TaskGetAllUserTask(c.GetInt("id"), (pageNum-1)*pageSize, pageSize, model.SyncTaskQueryParams{
Platform: constant.TaskPlatform(strconv.Itoa(constant.ChannelTypeKlingAiping)),
Action: route.Action,
})
data := make([]any, 0, len(tasks))
for _, task := range tasks {
data = append(data, taskDataObject(task))
}
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "success",
"request_id": c.GetString(common.RequestIdKey),
"data": data,
})
}

func KlingAipingNativeProxy(c *gin.Context) {
route, ok := klingaiping.FindRoute(c.Request.Method, c.FullPath(), klingaiping.RouteKindProxy)
if !ok {
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "route not found"})
return
}
group := common.GetContextKeyString(c, constant.ContextKeyUsingGroup)
channel := resolveKlingAipingBoundChannelForProxy(c, group, route.BillingModel)
if channel == nil {
var err error
channel, _, err = service.CacheGetRandomSatisfiedChannel(&service.RetryParam{
Ctx: c,
TokenGroup: group,
ModelName: route.BillingModel,
Retry: common.GetPointer(0),
AllowedChannelTypes: service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilyKling),
})
if err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"code": 503, "message": err.Error()})
return
}
}
if setupErr := middleware.SetupContextForSelectedChannel(c, channel, route.BillingModel); setupErr != nil {
c.JSON(setupErr.StatusCode, gin.H{"code": setupErr.GetErrorCode(), "message": setupErr.Error()})
return
}

resp, err := doKlingAipingProxyRequest(c, route, channel)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"code": 502, "message": err.Error()})
return
}
defer resp.Body.Close()
copyProxyResponse(c, resp)
}

func resolveKlingAipingBoundChannelForProxy(c *gin.Context, group, billingModel string) *model.Channel {
group = strings.TrimSpace(group)
if group == "" || group == "auto" || strings.TrimSpace(billingModel) == "" {
return nil
}
channel, err := service.GetBoundKlingAssetChannelForModel(c.GetInt("id"), group, billingModel)
if err != nil {
logger.LogError(c, fmt.Sprintf("resolve kling aiping bound channel (proxy) failed: %v", err))
return nil
}
if channel == nil || !service.IsUsableKlingAssetChannel(channel, group) {
return nil
}
if !model.IsChannelEnabledForGroupModel(group, billingModel, channel.Id) {
return nil
}
return channel
}

func readJSONPayload(c *gin.Context) (map[string]any, error) {
body, err := common.GetBodyStorage(c)
if err != nil {
return nil, err
}
data, err := body.Bytes()
if err != nil {
return nil, err
}
_, _ = body.Seek(0, io.SeekStart)
c.Request.Body = io.NopCloser(body)
payload := map[string]any{}
if strings.TrimSpace(string(data)) == "" {
return payload, nil
}
if err := common.Unmarshal(data, &payload); err != nil {
return nil, err
}
return payload, nil
}

func resolveKlingAipingModel(payload map[string]any, route klingaiping.Route) string {
if modelName := stringFromMap(payload, "model_name"); modelName != "" {
return modelName
}
if modelName := stringFromMap(payload, "model"); modelName != "" {
return modelName
}
if route.BillingModel != "" {
return route.BillingModel
}
return "kling-v3"
}

func stringFromMap(payload map[string]any, key string) string {
if value, ok := payload[key].(string); ok {
return strings.TrimSpace(value)
}
return ""
}

func getKlingAipingUserTask(c *gin.Context, taskID string, action string) (*model.Task, bool) {
task, exist, err := model.GetByTaskId(c.GetInt("id"), taskID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": err.Error()})
return nil, false
}
if !exist || task.Platform != constant.TaskPlatform(strconv.Itoa(constant.ChannelTypeKlingAiping)) || task.Action != action {
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "task not found"})
return nil, false
}
return task, true
}

func refreshKlingAipingTaskIfNeeded(task *model.Task) {
if task.Status == model.TaskStatusSuccess || task.Status == model.TaskStatusFailure {
return
}
channelModel, err := model.GetChannelById(task.ChannelId, true)
if err != nil || channelModel == nil {
return
}
adaptor := relay.GetTaskAdaptor(task.Platform)
if adaptor == nil {
return
}
resp, err := adaptor.FetchTask(channelModel.GetBaseURL(), channelModel.Key, map[string]any{
"task_id": task.GetUpstreamTaskID(),
"action": task.Action,
}, channelModel.GetSetting().Proxy)
if err != nil || resp == nil {
return
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil || resp.StatusCode < 200 || resp.StatusCode >= 300 {
return
}
taskInfo, err := adaptor.ParseTaskResult(body)
if err != nil || taskInfo == nil {
return
}
snap := task.Snapshot()
task.Data = body
if taskInfo.Status != "" {
task.Status = model.TaskStatus(taskInfo.Status)
}
if taskInfo.Progress != "" {
task.Progress = taskInfo.Progress
}
if taskInfo.Url != "" {
task.PrivateData.ResultURL = taskInfo.Url
}
if !snap.Equal(task.Snapshot()) {
_, _ = task.UpdateWithStatus(snap.Status)
}
}

func buildKlingAipingTaskPayload(task *model.Task) map[string]any {
payload := map[string]any{
"code": 0,
"message": "success",
"data": taskDataObject(task),
}
return payload
}

func taskDataObject(task *model.Task) map[string]any {
payload := map[string]any{}
_ = common.Unmarshal(task.Data, &payload)
delete(payload, "aiping_id")
data, _ := payload["data"].(map[string]any)
if data == nil {
data = map[string]any{}
}
data["task_id"] = task.TaskID
if _, ok := data["task_status"]; !ok {
data["task_status"] = mapKlingAipingTaskStatus(task.Status)
}
if _, ok := data["task_status_msg"]; !ok {
data["task_status_msg"] = task.FailReason
}
if _, ok := data["created_at"]; !ok && task.CreatedAt != 0 {
data["created_at"] = task.CreatedAt
}
if _, ok := data["updated_at"]; !ok && task.UpdatedAt != 0 {
data["updated_at"] = task.UpdatedAt
}
ensureKlingAipingWatermarkURL(data)
return data
}

func mapKlingAipingTaskStatus(status model.TaskStatus) string {
switch status {
case model.TaskStatusSubmitted, model.TaskStatusQueued:
return "submitted"
case model.TaskStatusInProgress:
return "processing"
case model.TaskStatusSuccess:
return "succeed"
case model.TaskStatusFailure:
return "failed"
default:
return "processing"
}
}

func ensureKlingAipingWatermarkURL(data map[string]any) {
taskResult, _ := data["task_result"].(map[string]any)
if taskResult == nil {
return
}
videos, _ := taskResult["videos"].([]any)
for _, videoAny := range videos {
video, _ := videoAny.(map[string]any)
if video == nil {
continue
}
if _, ok := video["watermark_url"]; !ok {
video["watermark_url"] = ""
}
}
}

func parseKlingAipingPage(c *gin.Context) (int, int, error) {
pageNum := parseIntDefault(c.Query("pageNum"), 1)
pageSize := parseIntDefault(c.Query("pageSize"), 30)
if pageNum < 1 || pageNum > 1000 {
return 0, 0, fmt.Errorf("pageNum must be in [1, 1000]")
}
if pageSize < 1 || pageSize > 500 {
return 0, 0, fmt.Errorf("pageSize must be in [1, 500]")
}
return pageNum, pageSize, nil
}

func parseIntDefault(raw string, fallback int) int {
if strings.TrimSpace(raw) == "" {
return fallback
}
v, err := strconv.Atoi(raw)
if err != nil {
return -1
}
return v
}

func doKlingAipingProxyRequest(c *gin.Context, route klingaiping.Route, channelModel *model.Channel) (*http.Response, error) {
baseURL := strings.TrimRight(channelModel.GetBaseURL(), "/")
upstreamPath := strings.Replace(route.UpstreamPath, ":id", c.Param("id"), 1)
url := baseURL + upstreamPath
if c.Request.URL.RawQuery != "" {
url += "?" + c.Request.URL.RawQuery
}
var body io.Reader
if c.Request.Method != http.MethodGet {
data, err := proxyBodyBytes(c, route)
if err != nil {
return nil, err
}
body = bytes.NewReader(data)
}
req, err := http.NewRequest(c.Request.Method, url, body)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/json")
req.Header.Set("Content-Type", "application/json")
key := common.GetContextKeyString(c, constant.ContextKeyChannelKey)
if key == "" {
key = channelModel.Key
}
req.Header.Set("Authorization", "Bearer "+key)
client, err := service.GetHttpClientWithProxy(channelModel.GetSetting().Proxy)
if err != nil {
return nil, err
}
return client.Do(req)
}

func proxyBodyBytes(c *gin.Context, route klingaiping.Route) ([]byte, error) {
storage, err := common.GetBodyStorage(c)
if err != nil {
return nil, err
}
return storage.Bytes()
}

func copyProxyResponse(c *gin.Context, resp *http.Response) {
if resp.StatusCode >= http.StatusBadRequest {
copyNormalizedProxyError(c, resp)
return
}
for key, values := range resp.Header {
for _, value := range values {
c.Writer.Header().Add(key, value)
}
}
c.Status(resp.StatusCode)
_, _ = io.Copy(c.Writer, resp.Body)
}

func copyNormalizedProxyError(c *gin.Context, resp *http.Response) {
body, _ := io.ReadAll(resp.Body)
message := strings.TrimSpace(string(body))
payload := map[string]any{}
if len(body) > 0 && common.Unmarshal(body, &payload) == nil {
if msg := stringFromMap(payload, "message"); msg != "" {
message = msg
} else if msg := stringFromMap(payload, "msg"); msg != "" {
message = msg
} else if detail, ok := payload["detail"]; ok {
if detailMap, ok := detail.(map[string]any); ok {
if msg := stringFromMap(detailMap, "message"); msg != "" {
message = msg
} else if msg := stringFromMap(detailMap, "msg"); msg != "" {
message = msg
} else {
message = fmt.Sprint(detail)
}
} else {
message = fmt.Sprint(detail)
}
}
delete(payload, "msg")
} else {
payload = map[string]any{}
}
if message == "" {
message = resp.Status
}
if _, ok := payload["code"]; !ok {
payload["code"] = resp.StatusCode
}
payload["message"] = message
payload["request_id"] = c.GetString(common.RequestIdKey)
c.JSON(resp.StatusCode, payload)
}

func normalizeKlingAipingTaskError(taskErr *dto.TaskError) {
if taskErr == nil || strings.TrimSpace(taskErr.Message) == "" {
return
}
payload := map[string]any{}
if common.Unmarshal([]byte(taskErr.Message), &payload) != nil {
return
}
if msg := stringFromMap(payload, "message"); msg != "" {
taskErr.Message = msg
return
}
if msg := stringFromMap(payload, "msg"); msg != "" {
taskErr.Message = msg
return
}
if detail, ok := payload["detail"].(map[string]any); ok {
if msg := stringFromMap(detail, "message"); msg != "" {
taskErr.Message = msg
}
}
}

+ 250
- 0
controller/kling_aiping_native_test.go View File

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

import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/model"
klingaiping "github.com/QuantumNous/new-api/relay/channel/task/kling/aiping"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestKlingAipingSubmitPreparationUsesModelNameBeforeModel(t *testing.T) {
c := newControllerJSONContext(t, "/v1/videos/text2video", `{
"model":"Kling-V1.6",
"model_name":"Kling-V2.6",
"prompt":"prompt"
}`)
payload, err := readJSONPayload(c)
require.NoError(t, err)
route, ok := klingaiping.FindRoute(http.MethodPost, "/v1/videos/text2video", klingaiping.RouteKindSubmit)
require.True(t, ok)
info := &relaycommon.RelayInfo{TaskRelayInfo: &relaycommon.TaskRelayInfo{}}
modelName := resolveKlingAipingModel(payload, route)
relaycommon.StoreTaskRequest(c, info, route.Action, relaycommon.TaskSubmitReq{
Model: modelName,
Prompt: stringFromMap(payload, "prompt"),
Metadata: payload,
})

stored, err := relaycommon.GetTaskRequest(c)
require.NoError(t, err)
require.Equal(t, "Kling-V2.6", stored.Model)
require.Equal(t, "prompt", stored.Prompt)
require.Equal(t, "Kling-V1.6", stored.Metadata["model"])
require.Equal(t, "Kling-V2.6", stored.Metadata["model_name"])
}

func TestKlingAipingSubmitPreparationDoesNotLockChannel(t *testing.T) {
c := newControllerJSONContext(t, "/v1/general/custom-voices", `{
"voice_url":"https://example.com/voice.mp3",
"voice_name":"voice"
}`)
payload, err := readJSONPayload(c)
require.NoError(t, err)
route, ok := klingaiping.FindRoute(http.MethodPost, "/v1/general/custom-voices", klingaiping.RouteKindSubmit)
require.True(t, ok)
info := &relaycommon.RelayInfo{TaskRelayInfo: &relaycommon.TaskRelayInfo{}}
info.OriginModelName = resolveKlingAipingModel(payload, route)
info.Action = route.Action
relaycommon.StoreTaskRequest(c, info, route.Action, relaycommon.TaskSubmitReq{
Model: info.OriginModelName,
Metadata: payload,
})

require.Equal(t, klingaiping.ModelCustomVoices, info.OriginModelName)
require.Equal(t, klingaiping.ActionVoicesCreate, info.Action)
require.Nil(t, info.LockedChannel)
}

func TestConfigureKlingAipingTaskRelayInfoForcesChannelSelection(t *testing.T) {
c := newControllerJSONContext(t, "/v1/videos/text2video", `{"model_name":"Kling-V2.6","prompt":"prompt"}`)
payload, err := readJSONPayload(c)
require.NoError(t, err)
route, ok := klingaiping.FindRoute(http.MethodPost, "/v1/videos/text2video", klingaiping.RouteKindSubmit)
require.True(t, ok)
info := &relaycommon.RelayInfo{TaskRelayInfo: &relaycommon.TaskRelayInfo{}}

configureKlingAipingTaskRelayInfo(c, info, route, payload)

require.NotNil(t, info.ChannelMeta)
require.Zero(t, info.ChannelMeta.ChannelType)
require.Equal(t, "Kling-V2.6", info.OriginModelName)
require.Equal(t, klingaiping.ActionText2Video, info.Action)
require.Nil(t, info.LockedChannel)
}

func TestConfigureKlingAipingTaskRelayInfoStoresDuration(t *testing.T) {
c := newControllerJSONContext(t, "/v1/videos/text2video", `{"model_name":"Kling-V2.6","prompt":"prompt","duration":5}`)
payload, err := readJSONPayload(c)
require.NoError(t, err)
route, ok := klingaiping.FindRoute(http.MethodPost, "/v1/videos/text2video", klingaiping.RouteKindSubmit)
require.True(t, ok)
info := &relaycommon.RelayInfo{TaskRelayInfo: &relaycommon.TaskRelayInfo{}}

configureKlingAipingTaskRelayInfo(c, info, route, payload)

stored, err := relaycommon.GetTaskRequest(c)
require.NoError(t, err)
require.Equal(t, 5, stored.Duration)
require.Empty(t, stored.Seconds)
}

func TestRequiredTaskChannelTypeOnlyMatchesKlingAipingNativePaths(t *testing.T) {
c := newControllerJSONContext(t, "/v1/videos/text2video", `{}`)
c.Request.URL.Path = "/v1/videos/text2video"
require.Zero(t, requiredTaskChannelTypeForRequest(c))
require.ElementsMatch(t,
service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilyKling),
allowedTaskChannelTypesForRequest(c),
)

c = newControllerJSONContext(t, "/v1/videos/video-extend", `{}`)
c.Request.URL.Path = "/v1/videos/video-extend"
require.Zero(t, requiredTaskChannelTypeForRequest(c))
require.ElementsMatch(t,
service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilyKling),
allowedTaskChannelTypesForRequest(c),
)

c = newControllerJSONContext(t, "/v1/videos/video_123/remix", `{}`)
c.Request.URL.Path = "/v1/videos/video_123/remix"
require.Zero(t, requiredTaskChannelTypeForRequest(c))
require.Nil(t, allowedTaskChannelTypesForRequest(c))
}

func TestKlingAipingTaskDataObjectUsesPublicTaskIDAndWatermarkURL(t *testing.T) {
task := &model.Task{
TaskID: "task_public",
Status: model.TaskStatusSuccess,
CreatedAt: 100,
UpdatedAt: 200,
Data: []byte(`{
"code":0,
"aiping_id":"internal",
"data":{
"task_id":"899333358055493641",
"task_status":"succeed",
"task_result":{"videos":[{"id":"v1","url":"https://example.com/video.mp4","duration":"5.041"}]}
}
}`),
}

data := taskDataObject(task)
require.Equal(t, "task_public", data["task_id"])
taskResult := data["task_result"].(map[string]any)
videos := taskResult["videos"].([]any)
require.Equal(t, "", videos[0].(map[string]any)["watermark_url"])
}

func TestDoKlingAipingProxyRequestUsesSelectedContextKey(t *testing.T) {
service.InitHttpClient()
var gotAuth string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotAuth = r.Header.Get("Authorization")
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":0,"message":"success","data":[]}`))
}))
defer server.Close()

c := newControllerJSONContext(t, "/v1/general/advanced-presets-elements", ``)
common.SetContextKey(c, constant.ContextKeyChannelKey, "selected-key")
route, ok := klingaiping.FindRoute(http.MethodGet, "/v1/general/advanced-presets-elements", klingaiping.RouteKindProxy)
require.True(t, ok)
channel := &model.Channel{
Key: "raw-channel-key",
BaseURL: common.GetPointer(server.URL),
}

resp, err := doKlingAipingProxyRequest(c, route, channel)
require.NoError(t, err)
defer resp.Body.Close()

require.Equal(t, "Bearer selected-key", gotAuth)
}

func TestCopyProxyResponseNormalizesMsgError(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Set(common.RequestIdKey, "req-test")
resp := &http.Response{
StatusCode: http.StatusUnauthorized,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"code":401,"msg":"unauthorized","data":null}`)),
}

copyProxyResponse(c, resp)

require.Equal(t, http.StatusUnauthorized, w.Code)
require.Contains(t, w.Body.String(), `"message":"unauthorized"`)
require.NotContains(t, w.Body.String(), `"msg"`)
require.Contains(t, w.Body.String(), `"request_id":"req-test"`)
}

func TestCopyProxyResponseNormalizesPlainTextError(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Set(common.RequestIdKey, "req-test")
resp := &http.Response{
StatusCode: http.StatusMethodNotAllowed,
Header: http.Header{"Content-Type": []string{"text/plain"}},
Body: io.NopCloser(strings.NewReader("Method Not Allowed")),
}

copyProxyResponse(c, resp)

require.Equal(t, http.StatusMethodNotAllowed, w.Code)
require.Contains(t, w.Body.String(), `"message":"Method Not Allowed"`)
require.Contains(t, w.Body.String(), `"request_id":"req-test"`)
}

func TestCopyProxyResponseNormalizesDetailMessageError(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Set(common.RequestIdKey, "req-test")
resp := &http.Response{
StatusCode: http.StatusServiceUnavailable,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"detail":{"message":"not found","error_type":"not_found"},"aiping_id":"internal"}`)),
}

copyProxyResponse(c, resp)

require.Equal(t, http.StatusServiceUnavailable, w.Code)
require.Contains(t, w.Body.String(), `"message":"not found"`)
require.NotContains(t, w.Body.String(), `map[`)
require.Contains(t, w.Body.String(), `"request_id":"req-test"`)
}

func TestNormalizeKlingAipingTaskErrorMessageExtractsUpstreamJSONMessage(t *testing.T) {
taskErr := &dto.TaskError{
Code: "fail_to_fetch_task",
Message: `{"code":400,"message":"ERROR: image download failed","request_id":"upstream"}`,
StatusCode: http.StatusBadRequest,
}

normalizeKlingAipingTaskError(taskErr)

require.Equal(t, "ERROR: image download failed", taskErr.Message)
}

func TestParseKlingAipingPageBoundaries(t *testing.T) {
c := newControllerJSONContext(t, "/v1/videos/text2video?pageNum=1001&pageSize=30", `{}`)
c.Request.URL.RawQuery = "pageNum=1001&pageSize=30"
_, _, err := parseKlingAipingPage(c)
require.ErrorContains(t, err, "pageNum")

c = newControllerJSONContext(t, "/v1/videos/text2video?pageNum=1&pageSize=501", `{}`)
c.Request.URL.RawQuery = "pageNum=1&pageSize=501"
_, _, err = parseKlingAipingPage(c)
require.ErrorContains(t, err, "pageSize")
}

+ 10
- 2
controller/log.go View File

@@ -21,7 +21,9 @@ func GetAllLogs(c *gin.Context) {
channel, _ := strconv.Atoi(c.Query("channel"))
group := c.Query("group")
requestId := c.Query("request_id")
logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId)
chatId := c.Query("chat_id")
upstreamId := c.Query("upstream_id")
logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId, chatId, upstreamId)
if err != nil {
common.ApiError(c, err)
return
@@ -42,11 +44,17 @@ func GetUserLogs(c *gin.Context) {
modelName := c.Query("model_name")
group := c.Query("group")
requestId := c.Query("request_id")
logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId)
chatId := c.Query("chat_id")
upstreamId := c.Query("upstream_id")
logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId, chatId, upstreamId)
if err != nil {
common.ApiError(c, err)
return
}
for i := range logs {
logs[i].ChatId = ""
logs[i].UpstreamId = ""
}
pageInfo.SetTotal(int(total))
pageInfo.SetItems(logs)
common.ApiSuccess(c, pageInfo)


+ 94
- 0
controller/log_identity_test.go View File

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

import (
"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/require"
"gorm.io/gorm"
)

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

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

origLogDB := model.LOG_DB
origSQLite := common.UsingSQLite
origMySQL := common.UsingMySQL
origPostgreSQL := common.UsingPostgreSQL

model.LOG_DB = db
common.UsingSQLite = true
common.UsingMySQL = false
common.UsingPostgreSQL = false

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

t.Cleanup(func() {
model.LOG_DB = origLogDB
common.UsingSQLite = origSQLite
common.UsingMySQL = origMySQL
common.UsingPostgreSQL = origPostgreSQL
})

return db
}

func TestGetAllLogsPassesChatIDAndUpstreamIDFilters(t *testing.T) {
db := setupControllerLogIdentityDB(t)

require.NoError(t, db.Create(&model.Log{
UserId: 1,
Username: "alice",
CreatedAt: 1714465001,
Type: model.LogTypeConsume,
Content: "usage-a",
ModelName: "gpt-4o-mini",
TokenName: "demo",
RequestId: "req_a",
ChatId: "chat_a",
UpstreamId: "up_a",
}).Error)
require.NoError(t, db.Create(&model.Log{
UserId: 1,
Username: "alice",
CreatedAt: 1714465002,
Type: model.LogTypeConsume,
Content: "usage-b",
ModelName: "gpt-4o-mini",
TokenName: "demo",
RequestId: "req_b",
ChatId: "chat_b",
UpstreamId: "up_b",
}).Error)

gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/api/log", GetAllLogs)

req := httptest.NewRequest(http.MethodGet, "/api/log?type=2&chat_id=chat_a&upstream_id=up_a", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

var payload struct {
Success bool `json:"success"`
Data struct {
Items []model.Log `json:"items"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &payload))
require.True(t, payload.Success)
require.Len(t, payload.Data.Items, 1)
require.Equal(t, "chat_a", payload.Data.Items[0].ChatId)
require.Equal(t, "up_a", payload.Data.Items[0].UpstreamId)
}

+ 74
- 5
controller/misc.go View File

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

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

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

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

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

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

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

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

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

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


+ 59
- 0
controller/model_display_pricing_controller.go View File

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

import (
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)

func GetModelDisplayPricingRules(c *gin.Context) {
common.ApiSuccess(c, ratio_setting.GetModelDisplayPricingCopy())
}

func GetModelDisplayPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
items := ratio_setting.GetModelDisplayPricing(modelName)
if len(items) == 0 {
c.JSON(http.StatusNotFound, gin.H{"success": false, "message": "model display pricing not found"})
return
}
common.ApiSuccess(c, items)
}

func UpdateModelDisplayPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
var items []types.ModelDisplayPricingItem
if err := common.DecodeJson(c.Request.Body, &items); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := ratio_setting.SetModelDisplayPricing(modelName, items); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := model.UpdateOption(ratio_setting.ModelDisplayPricingOptionKey, ratio_setting.ModelDisplayPricing2JSONString()); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
return
}
model.RefreshPricing()
common.ApiSuccess(c, ratio_setting.GetModelDisplayPricing(modelName))
}

func DeleteModelDisplayPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
if err := ratio_setting.DeleteModelDisplayPricing(modelName); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := model.UpdateOption(ratio_setting.ModelDisplayPricingOptionKey, ratio_setting.ModelDisplayPricing2JSONString()); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
return
}
model.RefreshPricing()
common.ApiSuccess(c, nil)
}

+ 96
- 0
controller/model_display_pricing_controller_test.go View File

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

import (
"net/http"
"testing"

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

func setupModelDisplayPricingControllerTest(t *testing.T) *gin.Engine {
t.Helper()
gin.SetMode(gin.TestMode)

db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Option{}))
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
origUsingSQLite := common.UsingSQLite
origRedisEnabled := common.RedisEnabled
model.DB = db
common.UsingSQLite = true
common.RedisEnabled = false
model.InitOptionMap()
require.NoError(t, ratio_setting.UpdateModelDisplayPricingByJSONString(`{}`))

t.Cleanup(func() {
model.DB = origDB
common.UsingSQLite = origUsingSQLite
common.RedisEnabled = origRedisEnabled
require.NoError(t, ratio_setting.UpdateModelDisplayPricingByJSONString(`{}`))
require.NoError(t, sqlDB.Close())
})

r := gin.New()
r.GET("/api/option/model_display_pricing", GetModelDisplayPricingRules)
r.GET("/api/option/model_display_pricing/*model", GetModelDisplayPricingRule)
r.PUT("/api/option/model_display_pricing/*model", UpdateModelDisplayPricingRule)
r.DELETE("/api/option/model_display_pricing/*model", DeleteModelDisplayPricingRule)
return r
}

func TestModelDisplayPricingControllerCRUDAndWildcard(t *testing.T) {
r := setupModelDisplayPricingControllerTest(t)
items := []types.ModelDisplayPricingItem{
{Specification: "768p-6s", OfficialSupplierTip: "768P 6s", Price: 0.33333, Unit: "second", SortOrder: 1},
}

w := performJSONRequest(t, r, http.MethodPut, "/api/option/model_display_pricing/provider/kling-video", items)
require.Equal(t, http.StatusOK, w.Code)
require.Len(t, ratio_setting.GetModelDisplayPricing("provider/kling-video"), 1)

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_display_pricing/provider/kling-video", nil)
require.Equal(t, http.StatusOK, w.Code)
require.Contains(t, w.Body.String(), "768p-6s")

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_display_pricing", nil)
require.Equal(t, http.StatusOK, w.Code)
require.Contains(t, w.Body.String(), "provider/kling-video")

w = performJSONRequest(t, r, http.MethodDelete, "/api/option/model_display_pricing/provider/kling-video", nil)
require.Equal(t, http.StatusOK, w.Code)
require.Empty(t, ratio_setting.GetModelDisplayPricing("provider/kling-video"))
}

func TestModelDisplayPricingControllerValidationError(t *testing.T) {
r := setupModelDisplayPricingControllerTest(t)
items := []types.ModelDisplayPricingItem{{Specification: "", Price: 1, Unit: "second"}}

w := performJSONRequest(t, r, http.MethodPut, "/api/option/model_display_pricing/kling-video", items)
require.Equal(t, http.StatusBadRequest, w.Code)
require.Contains(t, w.Body.String(), "specification is required")
}

func TestModelDisplayPricingOptionMapHydratesCache(t *testing.T) {
r := setupModelDisplayPricingControllerTest(t)
require.NotNil(t, r)

require.NoError(t, model.UpdateOption(ratio_setting.ModelDisplayPricingOptionKey, `{
"kling-video":[{"specification":"768p","price":1,"unit":"second"}]
}`))

items := ratio_setting.GetModelDisplayPricing("kling-video")
require.Len(t, items, 1)
require.Equal(t, "768p", items[0].Specification)
}

+ 8
- 0
controller/model_meta.go View File

@@ -97,6 +97,10 @@ func CreateModelMeta(c *gin.Context) {
return
}

if strings.TrimSpace(m.Endpoints) == "" {
m.Endpoints = `["openai"]`
}

if err := m.Insert(); err != nil {
common.ApiError(c, err)
return
@@ -135,6 +139,10 @@ func UpdateModelMeta(c *gin.Context) {
return
}

if strings.TrimSpace(m.Endpoints) == "" {
m.Endpoints = `["openai"]`
}

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


+ 89
- 0
controller/model_meta_test.go View File

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

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

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

func setupModelMetaTestDB(t *testing.T) {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Model{}))

origDB := model.DB
model.DB = db
t.Cleanup(func() {
model.RefreshPricing()
model.DB = origDB
})
}

func TestCreateModelMetaDefaultsEmptyEndpointsToOpenAI(t *testing.T) {
setupModelMetaTestDB(t)
gin.SetMode(gin.TestMode)
r := gin.New()
r.POST("/api/models/", CreateModelMeta)

req := httptest.NewRequest(
http.MethodPost,
"/api/models/",
bytes.NewBufferString(`{"model_name":"custom-chat-model"}`),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()

r.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var resp map[string]any
require.NoError(t, common.Unmarshal(w.Body.Bytes(), &resp))
require.True(t, resp["success"].(bool))

var saved model.Model
require.NoError(t, model.DB.Where("model_name = ?", "custom-chat-model").First(&saved).Error)
require.Equal(t, `["openai"]`, saved.Endpoints)
}

func TestUpdateModelMetaDefaultsEmptyEndpointsToOpenAI(t *testing.T) {
setupModelMetaTestDB(t)
gin.SetMode(gin.TestMode)
existing := &model.Model{
ModelName: "custom-chat-model",
Endpoints: `{"anthropic":{"path":"/v1/messages","method":"POST"}}`,
Status: 1,
}
require.NoError(t, existing.Insert())

r := gin.New()
r.PUT("/api/models/", UpdateModelMeta)

req := httptest.NewRequest(
http.MethodPut,
"/api/models/",
bytes.NewBufferString(`{"id":`+strconv.Itoa(existing.Id)+`,"model_name":"custom-chat-model","endpoints":"","status":1}`),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()

r.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var resp map[string]any
require.NoError(t, common.Unmarshal(w.Body.Bytes(), &resp))
require.True(t, resp["success"].(bool))

var saved model.Model
require.NoError(t, model.DB.First(&saved, existing.Id).Error)
require.Equal(t, `["openai"]`, saved.Endpoints)
}

+ 59
- 0
controller/model_pricing_controller.go View File

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

import (
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)

func GetModelPricingRules(c *gin.Context) {
common.ApiSuccess(c, ratio_setting.GetPricingConfigCopy())
}

func GetModelPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
cfg := ratio_setting.GetPricingConfig(modelName)
if cfg == nil {
c.JSON(http.StatusNotFound, gin.H{"success": false, "message": "model pricing rule not found"})
return
}
common.ApiSuccess(c, cfg)
}

func UpdateModelPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
var cfg types.PricingConfig
if err := common.DecodeJson(c.Request.Body, &cfg); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := ratio_setting.SetPricingConfig(modelName, &cfg); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := model.UpdateOption(ratio_setting.ModelPricingRulesOptionKey, ratio_setting.ModelPricingRules2JSONString()); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
return
}
model.RefreshPricing()
common.ApiSuccess(c, ratio_setting.GetPricingConfig(modelName))
}

func DeleteModelPricingRule(c *gin.Context) {
modelName := strings.TrimPrefix(c.Param("model"), "/")
if err := ratio_setting.DeletePricingConfig(modelName); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}
if err := model.UpdateOption(ratio_setting.ModelPricingRulesOptionKey, ratio_setting.ModelPricingRules2JSONString()); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
return
}
model.RefreshPricing()
common.ApiSuccess(c, nil)
}

+ 153
- 0
controller/model_pricing_controller_test.go View File

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

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

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

func setupModelPricingControllerTest(t *testing.T) *gin.Engine {
t.Helper()
gin.SetMode(gin.TestMode)

db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Option{}))
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
origUsingSQLite := common.UsingSQLite
origRedisEnabled := common.RedisEnabled
model.DB = db
common.UsingSQLite = true
common.RedisEnabled = false
model.InitOptionMap()
require.NoError(t, ratio_setting.UpdateModelPricingRulesByJSONString(`{}`))
require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{"video-test":0.25}`))

t.Cleanup(func() {
model.DB = origDB
common.UsingSQLite = origUsingSQLite
common.RedisEnabled = origRedisEnabled
require.NoError(t, ratio_setting.UpdateModelPricingRulesByJSONString(`{}`))
require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{}`))
require.NoError(t, sqlDB.Close())
})

r := gin.New()
r.GET("/api/option/model_pricing", GetModelPricingRules)
r.GET("/api/option/model_pricing/*model", GetModelPricingRule)
r.PUT("/api/option/model_pricing/*model", UpdateModelPricingRule)
r.DELETE("/api/option/model_pricing/*model", DeleteModelPricingRule)
return r
}

func TestModelPricingController_CRUD(t *testing.T) {
r := setupModelPricingControllerTest(t)
cfg := validControllerPricingConfig()

w := performJSONRequest(t, r, http.MethodPut, "/api/option/model_pricing/video-test", cfg)
require.Equal(t, http.StatusOK, w.Code)
assert.True(t, responseSuccess(t, w.Body.Bytes()))

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_pricing/video-test", nil)
require.Equal(t, http.StatusOK, w.Code)
assert.True(t, responseSuccess(t, w.Body.Bytes()))
assert.Contains(t, w.Body.String(), `"billing_unit":"per_call"`)

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_pricing", nil)
require.Equal(t, http.StatusOK, w.Code)
assert.True(t, responseSuccess(t, w.Body.Bytes()))
assert.Contains(t, w.Body.String(), "video-test")

w = performJSONRequest(t, r, http.MethodDelete, "/api/option/model_pricing/video-test", nil)
require.Equal(t, http.StatusOK, w.Code)
assert.True(t, responseSuccess(t, w.Body.Bytes()))

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_pricing/video-test", nil)
assert.Equal(t, http.StatusNotFound, w.Code)
legacyPrice, ok := ratio_setting.GetModelPrice("video-test", false)
require.True(t, ok)
assert.Equal(t, 0.25, legacyPrice)
}

func TestModelPricingController_PutRejectsBadConfig(t *testing.T) {
r := setupModelPricingControllerTest(t)
cfg := validControllerPricingConfig()
cfg.Table = nil

w := performJSONRequest(t, r, http.MethodPut, "/api/option/model_pricing/video-test", cfg)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "table must not be empty")
}

func TestModelPricingController_CRUDWithSlashModelName(t *testing.T) {
r := setupModelPricingControllerTest(t)
cfg := validControllerPricingConfig()
modelName := "openai/gpt-4o"

w := performJSONRequest(t, r, http.MethodPut, "/api/option/model_pricing/"+modelName, cfg)
require.Equal(t, http.StatusOK, w.Code)
require.NotNil(t, ratio_setting.GetPricingConfig(modelName))

w = performJSONRequest(t, r, http.MethodGet, "/api/option/model_pricing/"+modelName, nil)
require.Equal(t, http.StatusOK, w.Code)
assert.Contains(t, w.Body.String(), `"billing_unit":"per_call"`)

w = performJSONRequest(t, r, http.MethodDelete, "/api/option/model_pricing/"+modelName, nil)
require.Equal(t, http.StatusOK, w.Code)
assert.Nil(t, ratio_setting.GetPricingConfig(modelName))
}

func performJSONRequest(t *testing.T, r *gin.Engine, method, path string, body any) *httptest.ResponseRecorder {
t.Helper()
var payload []byte
if body != nil {
var err error
payload, err = common.Marshal(body)
require.NoError(t, err)
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
return w
}

func responseSuccess(t *testing.T, payload []byte) bool {
t.Helper()
var body struct {
Success bool `json:"success"`
}
require.NoError(t, common.Unmarshal(payload, &body))
return body.Success
}

func validControllerPricingConfig() types.PricingConfig {
return types.PricingConfig{
SchemaVersion: 1,
Scope: types.PricingScopeModel,
BillingUnit: types.BillingUnitPerCall,
PreconsumeStrategy: types.PreconsumeStrategyExact,
Dimensions: []types.PricingDimension{
{Key: "resolution", Source: "request.resolution", Type: "string"},
},
Table: []types.PricingRow{
{"resolution": "720P", "price": 0.1, "source": types.PricingRowSourceManual},
},
Fallback: types.PricingFallback{Strategy: types.PricingFallbackReject},
}
}

+ 10
- 0
controller/playground.go View File

@@ -4,9 +4,11 @@ import (
"errors"
"fmt"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/middleware"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/types"

"github.com/gin-gonic/gin"
@@ -54,3 +56,11 @@ func Playground(c *gin.Context) {

Relay(c, types.RelayFormatOpenAI)
}

// GetPlaygroundConfig 获取 Playground 公开配置(无需认证)
func GetPlaygroundConfig(c *gin.Context) {
setting := operation_setting.GetPlaygroundSetting()
common.ApiSuccess(c, gin.H{
"mutual_exclusive_params": setting.MutualExclusiveParams,
})
}

+ 171
- 0
controller/pricing.go View File

@@ -1,14 +1,39 @@
package controller

import (
"fmt"
"math"
"sort"
"strings"

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

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

// resolveCurrentUser 从 gin.Context 解析当前登录用户信息。
// 未登录或查询失败时 loggedIn=false,groupRatio=1.0。
func resolveCurrentUser(c *gin.Context) (loggedIn bool, userId int, userGroup string, groupRatio float64) {
groupRatio = 1.0
raw, exists := c.Get("id")
if !exists {
return
}
userId = raw.(int)
user, err := model.GetUserCache(userId)
if err != nil {
return
}
loggedIn = true
userGroup = user.Group
groupRatio = service.GetUserGroupRatio(userGroup, userGroup)
return
}

func filterPricingByUsableGroups(pricing []model.Pricing, usableGroup map[string]string) []model.Pricing {
if len(pricing) == 0 {
return pricing
@@ -76,6 +101,152 @@ func GetPricing(c *gin.Context) {
})
}

// GetUserPricing 获取用户对指定模型的价格信息(原价 vs 用户价)
func GetUserPricing(c *gin.Context) {
modelName := c.Param("model")
// *model 通配符返回值带前导 /,需要去掉
modelName = strings.TrimPrefix(modelName, "/")
if modelName == "" {
common.ApiErrorMsg(c, "模型名不能为空")
return
}

pricingData := model.GetPricingByModel(modelName)
if pricingData == nil {
common.ApiErrorMsg(c, "未找到该模型的定价信息")
return
}

loggedIn, _, userGroup, groupRatio := resolveCurrentUser(c)
if !loggedIn {
respondOriginalPrice(c, pricingData)
return
}

totalRatio := groupRatio
savingsPercent := int(math.Round((1 - totalRatio) * 100))

result := gin.H{
"success": true,
"model_name": pricingData.ModelName,
"quota_type": pricingData.QuotaType,
"group": userGroup,
"group_ratio": groupRatio,
"savings_percent": savingsPercent,
"logged_in": true,
}
if pricingData.PricingConfig != nil {
result["pricing_config"] = pricingData.PricingConfig
}
if len(pricingData.DisplayPricing) > 0 {
result["display_pricing"] = pricingData.DisplayPricing
}

if pricingData.QuotaType == model.QuotaTypeByTokens {
if pricingData.PricingConfig != nil {
originalPrice := representativeMatrixPrice(pricingData.PricingConfig, pricingData.ModelPrice)
result["original_price"] = originalPrice
result["user_price"] = originalPrice * totalRatio
}
originalInput := pricingData.ModelRatio * 2
originalOutput := pricingData.ModelRatio * pricingData.CompletionRatio * 2
result["original_input"] = originalInput
result["original_output"] = originalOutput
result["user_input"] = originalInput * totalRatio
result["user_output"] = originalOutput * totalRatio
} else {
result["original_price"] = pricingData.ModelPrice
result["user_price"] = pricingData.ModelPrice * totalRatio
}

if savingsPercent > 0 {
result["discount"] = formatDiscount(totalRatio)
}

c.JSON(200, result)
}

// respondOriginalPrice 未登录或用户查询失败时,只返回原价
func respondOriginalPrice(c *gin.Context, pricingData *model.Pricing) {
result := gin.H{
"success": true,
"model_name": pricingData.ModelName,
"quota_type": pricingData.QuotaType,
"logged_in": false,
}
if pricingData.QuotaType == model.QuotaTypeByTokens {
if pricingData.PricingConfig != nil {
result["pricing_config"] = pricingData.PricingConfig
result["original_price"] = representativeMatrixPrice(pricingData.PricingConfig, pricingData.ModelPrice)
}
result["original_input"] = pricingData.ModelRatio * 2
result["original_output"] = pricingData.ModelRatio * pricingData.CompletionRatio * 2
} else {
if pricingData.PricingConfig != nil {
result["pricing_config"] = pricingData.PricingConfig
}
result["original_price"] = pricingData.ModelPrice
}
if len(pricingData.DisplayPricing) > 0 {
result["display_pricing"] = pricingData.DisplayPricing
}
c.JSON(200, result)
}

func representativeMatrixPrice(cfg *types.PricingConfig, fallback float64) float64 {
if cfg == nil || len(cfg.Table) == 0 {
return fallback
}
prices := make([]float64, 0, len(cfg.Table))
for _, row := range cfg.Table {
price, ok := pricingRowFloat(row["price"])
if ok {
prices = append(prices, price)
}
}
if len(prices) == 0 {
return fallback
}
sort.Float64s(prices)
mid := len(prices) / 2
if len(prices)%2 == 1 {
return prices[mid]
}
return (prices[mid-1] + prices[mid]) / 2
}

func pricingRowFloat(value any) (float64, bool) {
switch n := value.(type) {
case float64:
return n, true
case float32:
return float64(n), true
case int:
return float64(n), true
case int64:
return float64(n), true
case int32:
return float64(n), true
default:
return 0, false
}
}

// formatDiscount 将倍率转换为中文折扣格式
func formatDiscount(ratio float64) string {
if ratio <= 0 {
return "免费"
}
rawDiscount := math.Round(ratio*100) / 10
if rawDiscount >= 10 {
return ""
}
if rawDiscount == math.Trunc(rawDiscount) {
return fmt.Sprintf("%.0f折", rawDiscount)
}
return fmt.Sprintf("%.1f折", rawDiscount)
}

func ResetModelRatio(c *gin.Context) {
defaultStr := ratio_setting.DefaultModelRatio2JSONString()
err := model.UpdateOption("ModelRatio", defaultStr)


+ 0
- 131
controller/pricing_tag.go View File

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

import (
"strconv"

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

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

// GetAllPricingTags 获取所有定价标签
func GetAllPricingTags(c *gin.Context) {
list, err := model.GetAllPricingTags()
if err != nil {
common.ApiError(c, err)
return
}
common.ApiSuccess(c, list)
}

// CreatePricingTagRequest 创建定价标签请求
type CreatePricingTagRequest struct {
Name string `json:"name" binding:"required"`
Color string `json:"color"`
Description string `json:"description"`
SortOrder int `json:"sort_order"`
}

// CreatePricingTag 创建定价标签
func CreatePricingTag(c *gin.Context) {
var req CreatePricingTagRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiError(c, err)
return
}

// 检查名称是否重复
existing, _ := model.GetPricingTagByName(req.Name)
if existing != nil {
common.ApiErrorMsg(c, "tag name already exists")
return
}

pt := &model.PricingTag{
Name: req.Name,
Color: req.Color,
Description: req.Description,
SortOrder: req.SortOrder,
}
if pt.Color == "" {
pt.Color = "#1890ff"
}
if err := pt.Insert(); err != nil {
common.ApiError(c, err)
return
}

common.ApiSuccess(c, pt)
}

// UpdatePricingTagRequest 更新定价标签请求
type UpdatePricingTagRequest struct {
Name string `json:"name"`
Color string `json:"color"`
Description string `json:"description"`
SortOrder int `json:"sort_order"`
}

// UpdatePricingTag 更新定价标签
func UpdatePricingTag(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.Atoi(idStr)
if err != nil {
common.ApiError(c, err)
return
}

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

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

// 检查名称是否与其他标签重复
if req.Name != "" && req.Name != pt.Name {
existing, _ := model.GetPricingTagByName(req.Name)
if existing != nil && existing.Id != id {
common.ApiErrorMsg(c, "tag name already exists")
return
}
pt.Name = req.Name
}

if req.Color != "" {
pt.Color = req.Color
}
pt.Description = req.Description
pt.SortOrder = req.SortOrder

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

common.ApiSuccess(c, pt)
}

// DeletePricingTag 删除定价标签
func DeletePricingTag(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.Atoi(idStr)
if err != nil {
common.ApiError(c, err)
return
}

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

common.ApiSuccess(c, nil)
}

+ 349
- 0
controller/pricing_user_test.go View File

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

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

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

// setupPricingTestDB 初始化测试数据库
func setupPricingTestDB(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.User{}))

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

// setupPricingTestRouter 创建无认证路由(未登录场景)
func setupPricingTestRouter() *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/api/pricing/user/:model", GetUserPricing)
return r
}

// setupAuthRouter 创建带用户 ID 注入的路由(已登录场景)
func setupAuthRouter(userID int) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/api/pricing/user/:model", func(c *gin.Context) {
c.Set("id", userID)
c.Next()
}, GetUserPricing)
return r
}

// setPricingCache 直接设置定价缓存用于测试
func setPricingCache(pricing []model.Pricing) {
model.SetTestPricing(pricing)
}

// withGroupRatio 临时设置分组倍率,测试结束后恢复
func withGroupRatio(t *testing.T, jsonStr string) {
t.Helper()
original := ratio_setting.GetGroupRatioCopy()
ratio_setting.UpdateGroupRatioByJSONString(jsonStr)
t.Cleanup(func() {
origJSON, _ := json.Marshal(original)
ratio_setting.UpdateGroupRatioByJSONString(string(origJSON))
})
}

// ---- 测试用例 ----

// TestGetUserPricing_ModelNotFound 模型不存在时应返回错误
func TestGetUserPricing_ModelNotFound(t *testing.T) {
setupPricingTestDB(t)
router := setupPricingTestRouter()
setPricingCache([]model.Pricing{})

w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/nonexistent-model", nil)
router.ServeHTTP(w, req)

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

// TestGetUserPricing_NotLoggedIn 未登录用户应只返回原价
func TestGetUserPricing_NotLoggedIn(t *testing.T) {
setupPricingTestDB(t)
router := setupPricingTestRouter()

setPricingCache([]model.Pricing{
{
ModelName: "gpt-4o",
QuotaType: 0,
ModelRatio: 15,
CompletionRatio: 4,
EnableGroup: []string{"default", "vip"},
},
})

w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/gpt-4o", 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))
assert.Equal(t, false, resp["logged_in"])
assert.Equal(t, "gpt-4o", resp["model_name"])
assert.Equal(t, float64(0), resp["quota_type"])

// 验证原价: model_ratio * 2 = 15 * 2 = 30
assert.Equal(t, float64(30), resp["original_input"])
// 输出原价: model_ratio * completion_ratio * 2 = 15 * 4 * 2 = 120
assert.Equal(t, float64(120), resp["original_output"])

// 不应有用户价字段
_, hasUserInput := resp["user_input"]
assert.False(t, hasUserInput)
}

// TestGetUserPricing_NotLoggedIn_PerCall 按次计费模型,未登录
func TestGetUserPricing_NotLoggedIn_PerCall(t *testing.T) {
setupPricingTestDB(t)
router := setupPricingTestRouter()

setPricingCache([]model.Pricing{
{
ModelName: "dall-e-3",
QuotaType: 1,
ModelPrice: 0.04,
EnableGroup: []string{"default"},
},
})

w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/dall-e-3", 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))
assert.Equal(t, false, resp["logged_in"])
assert.Equal(t, float64(0.04), resp["original_price"])
}

func TestGetUserPricingIncludesDisplayPricingWithoutChangingComputedPrice(t *testing.T) {
setupPricingTestDB(t)
modelName := "display-pricing-user-model"
router := gin.New()
router.GET("/api/pricing/user/*model", GetUserPricing)

setPricingCache([]model.Pricing{
{
ModelName: modelName,
QuotaType: model.QuotaTypeByCall,
ModelPrice: 0.5,
EnableGroup: []string{"default"},
DisplayPricing: []types.ModelDisplayPricingItem{
{Specification: "768p-6s", Price: 0.33333, Unit: "second", SortOrder: 1, DiscountRate: 1},
},
},
})

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

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.Equal(t, http.StatusOK, w.Code)
require.Contains(t, resp, "display_pricing")
require.Equal(t, float64(0.5), resp["original_price"])
require.NotEqual(t, float64(0.33333), resp["original_price"])
}


func TestGetPricingIncludesDisplayPricingForModelSquare(t *testing.T) {
setupPricingTestDB(t)
modelName := "display-pricing-square-model"
router := gin.New()
router.GET("/api/pricing", GetPricing)

setPricingCache([]model.Pricing{
{
ModelName: modelName,
QuotaType: model.QuotaTypeByCall,
ModelPrice: 0.5,
EnableGroup: []string{"default"},
DisplayPricing: []types.ModelDisplayPricingItem{
{Specification: "768p-6s", Price: 0.33333, Unit: "second", SortOrder: 1, DiscountRate: 1},
},
},
})
withGroupRatio(t, `{"default":1}`)

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

var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.Equal(t, http.StatusOK, w.Code)
require.Equal(t, true, resp["success"])
data := resp["data"].([]interface{})
require.Len(t, data, 1)
item := data[0].(map[string]interface{})
require.Equal(t, modelName, item["model_name"])
require.Contains(t, item, "display_pricing")
displayPricing := item["display_pricing"].([]interface{})
require.Len(t, displayPricing, 1)
require.Equal(t, "768p-6s", displayPricing[0].(map[string]interface{})["specification"])
}

// TestGetUserPricing_LoggedIn_NoDiscount 已登录但无折扣(分组倍率=1,无个人倍率)
func TestGetUserPricing_LoggedIn_NoDiscount(t *testing.T) {
db := setupPricingTestDB(t)

setPricingCache([]model.Pricing{
{
ModelName: "gpt-4o",
QuotaType: 0,
ModelRatio: 15,
CompletionRatio: 4,
EnableGroup: []string{"default", "vip"},
},
})

// 创建测试用户(default 分组,默认倍率为 1)
user := &model.User{Id: 100, Group: "default", Username: "testuser", Status: 1}
require.NoError(t, db.Create(user).Error)

router := setupAuthRouter(100)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/gpt-4o", 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))
assert.Equal(t, true, resp["logged_in"])
// default 分组默认倍率为 1,无折扣
assert.Equal(t, float64(0), resp["savings_percent"])
}

// TestGetUserPricing_LoggedIn_GroupDiscount 已登录,有分组折扣
func TestGetUserPricing_LoggedIn_GroupDiscount(t *testing.T) {
db := setupPricingTestDB(t)

setPricingCache([]model.Pricing{
{
ModelName: "gpt-4o",
QuotaType: 0,
ModelRatio: 15,
CompletionRatio: 4,
EnableGroup: []string{"default", "vip"},
},
})

// 创建 VIP 用户
user := &model.User{Id: 200, Group: "vip", Username: "vipuser", Status: 1}
require.NoError(t, db.Create(user).Error)

// 设置 VIP 分组倍率为 0.8
withGroupRatio(t, `{"default":1,"vip":0.8}`)

router := setupAuthRouter(200)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/gpt-4o", 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))
assert.Equal(t, true, resp["logged_in"])
assert.Equal(t, "vip", resp["group"])
assert.Equal(t, float64(0.8), resp["group_ratio"])

// 用户价 = 原价 * 0.8
// 输入: 30 * 0.8 = 24
assert.Equal(t, float64(24), resp["user_input"])
// 输出: 120 * 0.8 = 96
assert.Equal(t, float64(96), resp["user_output"])

assert.Equal(t, float64(20), resp["savings_percent"])
assert.Equal(t, "8折", resp["discount"])
}

// TestGetUserPricing_PerCall_WithDiscount 按次计费 + 折扣
func TestGetUserPricing_PerCall_WithDiscount(t *testing.T) {
db := setupPricingTestDB(t)

setPricingCache([]model.Pricing{
{
ModelName: "dall-e-3",
QuotaType: 1,
ModelPrice: 0.04,
EnableGroup: []string{"default", "vip"},
},
})

user := &model.User{Id: 500, Group: "vip", Username: "percallvip", Status: 1}
require.NoError(t, db.Create(user).Error)

withGroupRatio(t, `{"default":1,"vip":0.5}`)

router := setupAuthRouter(500)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/pricing/user/dall-e-3", 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))
assert.Equal(t, float64(0.04), resp["original_price"])
assert.Equal(t, float64(0.02), resp["user_price"])
assert.Equal(t, float64(50), resp["savings_percent"])
assert.Equal(t, "5折", resp["discount"])
}

// TestFormatDiscount 折扣格式化测试
func TestFormatDiscount(t *testing.T) {
tests := []struct {
ratio float64
expected string
}{
{0.5, "5折"},
{0.8, "8折"},
{0.9, "9折"},
{0.85, "8.5折"},
{0.75, "7.5折"},
{0.95, "9.5折"},
{1.0, ""},
{0.0, "免费"},
}
for _, tt := range tests {
result := formatDiscount(tt.ratio)
assert.Equal(t, tt.expected, result, "ratio=%.2f", tt.ratio)
}
}

+ 2
- 0
controller/redemption.go View File

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


+ 159
- 0
controller/redemption_test.go View File

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

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

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

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

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

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

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

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

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

return db
}

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

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

return r
}

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

+ 229
- 38
controller/relay.go View File

@@ -6,10 +6,12 @@ import (
"io"
"log"
"net/http"
"strconv"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/common/metrics"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/logger"
@@ -22,6 +24,7 @@ import (
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/QuantumNous/new-api/types"

"github.com/bytedance/gopkg/util/gopool"
@@ -87,6 +90,15 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
defer func() {
if newAPIError != nil {
logger.LogError(c, fmt.Sprintf("relay error: %s", newAPIError.Error()))

// Prometheus: 错误计数
if metrics.IsEnabled() {
channelId := strconv.Itoa(c.GetInt("channel_id"))
modelName := c.GetString("original_model")
errorType, errorCode := metrics.ClassifyError(newAPIError.StatusCode, newAPIError.Error())
metrics.GetMetrics().RequestErrorsTotal.WithLabelValues(channelId, modelName, errorType, errorCode).Inc()
}

newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId))
switch relayFormat {
case types.RelayFormatOpenAIRealtime:
@@ -179,23 +191,21 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
}()

retryParam := &service.RetryParam{
Ctx: c,
TokenGroup: relayInfo.TokenGroup,
ModelName: relayInfo.OriginModelName,
Retry: common.GetPointer(0),
Ctx: c,
TokenGroup: relayInfo.TokenGroup,
ModelName: relayInfo.OriginModelName,
Retry: common.GetPointer(0),
RequireMatrixUsageBilling: relayInfo.RequireMatrixUsageBilling,
}

for ; retryParam.GetRetry() <= common.RetryTimes; retryParam.IncreaseRetry() {
channel, channelErr := getChannel(c, relayInfo, retryParam)
channel, _, channelErr := getChannel(c, relayInfo, retryParam)
if channelErr != nil {
logger.LogError(c, channelErr.Error())
newAPIError = channelErr
break
}

// 在渠道选择后更新价格数据以使用渠道定价
helper.UpdatePriceDataForChannelPricing(c, relayInfo, channel.Id)

addUsedChannel(c, channel.Id)
bodyStorage, bodyErr := common.GetBodyStorage(c)
if bodyErr != nil {
@@ -231,6 +241,12 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
if !shouldRetry(c, newAPIError, common.RetryTimes-retryParam.GetRetry()) {
break
}

// Prometheus: 重试计数
if metrics.IsEnabled() {
channelIdStr := strconv.Itoa(channel.Id)
metrics.GetMetrics().RequestRetriesTotal.WithLabelValues(channelIdStr, relayInfo.OriginModelName).Inc()
}
}

useChannel := c.GetStringSlice("use_channel")
@@ -280,7 +296,37 @@ func fastTokenCountMetaForPricing(request dto.Request) *types.TokenCountMeta {
return meta
}

func getChannel(c *gin.Context, info *relaycommon.RelayInfo, retryParam *service.RetryParam) (*model.Channel, *types.NewAPIError) {
func concreteTaskVideoBindingGroup(c *gin.Context, group string) string {
group = strings.TrimSpace(group)
if group == "" {
group = strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyTokenGroup))
}
if group == "auto" {
autoGroup := strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyAutoGroup))
if autoGroup == "" || autoGroup == "auto" {
return ""
}
return autoGroup
}
return group
}

func persistTaskVideoBindingIfNeeded(userId int, group string, channel *model.Channel) error {
group = strings.TrimSpace(group)
if group == "" || group == "auto" || channel == nil {
return nil
}
family, ok := service.VideoAssetFamilyForChannelType(channel.Type)
if !ok {
return nil
}
if !service.IsUsableVideoAssetChannelForFamily(channel, group, "", family) {
return nil
}
return service.BindVideoAssetChannel(userId, group, channel, family)
}

func getChannel(c *gin.Context, info *relaycommon.RelayInfo, retryParam *service.RetryParam) (*model.Channel, string, *types.NewAPIError) {
if info.ChannelMeta == nil {
autoBan := c.GetBool("auto_ban")
autoBanInt := 1
@@ -292,24 +338,41 @@ func getChannel(c *gin.Context, info *relaycommon.RelayInfo, retryParam *service
Type: c.GetInt("channel_type"),
Name: c.GetString("channel_name"),
AutoBan: &autoBanInt,
}, nil
}, concreteTaskVideoBindingGroup(c, info.TokenGroup), nil
}
channel, selectGroup, err := service.CacheGetRandomSatisfiedChannel(retryParam)

info.PriceData.GroupRatioInfo = helper.HandleGroupRatio(c, info)

if err != nil {
return nil, types.NewError(fmt.Errorf("获取分组 %s 下模型 %s 的可用渠道失败(retry): %s", selectGroup, info.OriginModelName, err.Error()), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry())
if apiErr, ok := err.(*types.NewAPIError); ok {
return nil, selectGroup, apiErr
}
return nil, selectGroup, types.NewError(fmt.Errorf("获取分组 %s 下模型 %s 的可用渠道失败(retry): %s", selectGroup, info.OriginModelName, err.Error()), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry())
}
if channel == nil {
return nil, types.NewError(fmt.Errorf("分组 %s 下模型 %s 的可用渠道不存在(retry)", selectGroup, info.OriginModelName), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry())
return nil, selectGroup, types.NewError(fmt.Errorf("分组 %s 下模型 %s 的可用渠道不存在(retry)", selectGroup, info.OriginModelName), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry())
}

newAPIError := middleware.SetupContextForSelectedChannel(c, channel, info.OriginModelName)
if newAPIError != nil {
return nil, newAPIError
return nil, selectGroup, newAPIError
}
return channel, nil
return channel, selectGroup, nil
}

func requiredTaskChannelTypeForRequest(c *gin.Context) int {
return 0
}

func allowedTaskChannelTypesForRequest(c *gin.Context) []int {
if strings.HasPrefix(c.Request.URL.Path, "/api/v3/contents/generations/tasks") {
return service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilySeedance)
}
if isKlingAipingNativePath(c.Request.URL.Path) {
return service.VideoAssetChannelTypesForFamily(service.VideoAssetFamilyKling)
}
return nil
}

func shouldRetry(c *gin.Context, openaiErr *types.NewAPIError, retryTimes int) bool {
@@ -328,9 +391,6 @@ func shouldRetry(c *gin.Context, openaiErr *types.NewAPIError, retryTimes int) b
if retryTimes <= 0 {
return false
}
if _, ok := c.Get("specific_channel_id"); ok {
return false
}
code := openaiErr.StatusCode
if code >= 200 && code < 300 {
return false
@@ -369,6 +429,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)
@@ -470,6 +536,101 @@ func RelayTaskFetch(c *gin.Context) {
}
}

func preloadTaskPricingConfig(c *gin.Context, info *relaycommon.RelayInfo) *dto.TaskError {
modelName := strings.TrimSpace(info.OriginModelName)
action := strings.TrimSpace(info.Action)
contentType := c.Request.Header.Get("Content-Type")

storage, err := common.GetBodyStorage(c)
if err != nil {
status := http.StatusBadRequest
if common.IsRequestBodyTooLargeError(err) || errors.Is(err, common.ErrRequestBodyTooLarge) {
status = http.StatusRequestEntityTooLarge
}
return service.TaskErrorWrapperLocal(err, "read_request_body_failed", status)
}
defer func() {
_, _ = storage.Seek(0, io.SeekStart)
c.Request.Body = io.NopCloser(storage)
}()

switch {
case strings.HasPrefix(contentType, "application/json"):
body, err := storage.Bytes()
if err != nil {
return service.TaskErrorWrapperLocal(err, "read_request_body_failed", http.StatusBadRequest)
}
if strings.TrimSpace(string(body)) != "" {
var payload map[string]any
if err := common.Unmarshal(body, &payload); err != nil {
return service.TaskErrorWrapperLocal(err, "invalid_json", http.StatusBadRequest)
}
if raw, ok := payload["model_name"].(string); ok && strings.TrimSpace(raw) != "" {
modelName = strings.TrimSpace(raw)
} else if raw, ok := payload["model"].(string); ok && strings.TrimSpace(raw) != "" {
modelName = strings.TrimSpace(raw)
}
if raw, ok := payload["action"].(string); ok && strings.TrimSpace(raw) != "" {
action = strings.TrimSpace(raw)
}
}
case strings.Contains(contentType, gin.MIMEMultipartPOSTForm):
form, err := common.ParseMultipartFormReusable(c)
if err != nil {
return service.TaskErrorWrapperLocal(err, "invalid_multipart_form", http.StatusBadRequest)
}
if vals := form.Value["model"]; len(vals) > 0 && strings.TrimSpace(vals[0]) != "" {
modelName = strings.TrimSpace(vals[0])
}
if vals := form.Value["action"]; len(vals) > 0 && strings.TrimSpace(vals[0]) != "" {
action = strings.TrimSpace(vals[0])
}
}

if modelName == "" && action != "" {
platform := constant.TaskPlatform(c.GetString("platform"))
modelName = service.CoverTaskActionToModelName(platform, action)
}
if modelName != "" {
info.OriginModelName = modelName
}
if action != "" {
info.Action = action
}

info.PricingConfigSnapshotLoaded = true
info.PricingConfigSnapshot = ratio_setting.GetPricingConfig(info.OriginModelName)
info.RequireMatrixUsageBilling = info.PricingConfigSnapshot != nil &&
info.PricingConfigSnapshot.BillingUnit == types.BillingUnitPer1MTokens
return nil
}

func buildTaskBillingContext(info *relaycommon.RelayInfo) *model.TaskBillingContext {
bc := &model.TaskBillingContext{
ModelPrice: info.PriceData.ModelPrice,
GroupRatio: info.PriceData.GroupRatioInfo.GroupRatio,
ModelRatio: info.PriceData.ModelRatio,
OtherRatios: info.PriceData.OtherRatios,
OriginModelName: info.OriginModelName,
PerCallBilling: common.StringsContains(constant.TaskPricePatches, info.OriginModelName),
}
if decision := info.PricingDecisionFrozen; decision != nil && decision.BillingMode == types.BillingModeMatrix {
bc.BillingMode = decision.BillingMode
bc.PricingSnapshot = types.CloneMapAny(decision.Snapshot)
bc.BillingUnit = decision.BillingUnit
bc.TokenUnitPriceUSD = decision.TokenUnitPriceUSD
bc.PerCallBilling = decision.PerCallBilling
if decision.BillingUnit == types.BillingUnitPer1MTokens {
bc.ModelPrice = decision.TokenUnitPriceUSD
} else {
bc.ModelPrice = decision.PriceUSD
}
bc.GroupRatio = decision.GroupRatioInfo.GroupRatio
bc.OtherRatios = types.CloneRatios(decision.OtherRatios)
}
return bc
}

func RelayTask(c *gin.Context) {
relayInfo, err := relaycommon.GenRelayInfo(c, types.RelayFormatTask, nil, nil)
if err != nil {
@@ -480,11 +641,18 @@ func RelayTask(c *gin.Context) {
})
return
}
relayTaskWithInfo(c, relayInfo)
}

func relayTaskWithInfo(c *gin.Context, relayInfo *relaycommon.RelayInfo) {
if taskErr := relay.ResolveOriginTask(c, relayInfo); taskErr != nil {
respondTaskError(c, taskErr)
return
}
if taskErr := preloadTaskPricingConfig(c, relayInfo); taskErr != nil {
respondTaskError(c, taskErr)
return
}

var result *relay.TaskSubmitResult
var taskErr *dto.TaskError
@@ -495,26 +663,28 @@ func RelayTask(c *gin.Context) {
}()

retryParam := &service.RetryParam{
Ctx: c,
TokenGroup: relayInfo.TokenGroup,
ModelName: relayInfo.OriginModelName,
Retry: common.GetPointer(0),
Ctx: c,
TokenGroup: relayInfo.TokenGroup,
ModelName: relayInfo.OriginModelName,
Retry: common.GetPointer(0),
RequiredChannelType: requiredTaskChannelTypeForRequest(c),
AllowedChannelTypes: allowedTaskChannelTypesForRequest(c),
}

for ; retryParam.GetRetry() <= common.RetryTimes; retryParam.IncreaseRetry() {
var channel *model.Channel
var selectedGroup string

if lockedCh, ok := relayInfo.LockedChannel.(*model.Channel); ok && lockedCh != nil {
channel = lockedCh
if retryParam.GetRetry() > 0 {
if setupErr := middleware.SetupContextForSelectedChannel(c, channel, relayInfo.OriginModelName); setupErr != nil {
taskErr = service.TaskErrorWrapperLocal(setupErr.Err, "setup_locked_channel_failed", http.StatusInternalServerError)
break
}
selectedGroup = concreteTaskVideoBindingGroup(c, relayInfo.TokenGroup)
if setupErr := middleware.SetupContextForSelectedChannel(c, channel, relayInfo.OriginModelName); setupErr != nil {
taskErr = service.TaskErrorWrapperLocal(setupErr.Err, "setup_locked_channel_failed", http.StatusInternalServerError)
break
}
} else {
var channelErr *types.NewAPIError
channel, channelErr = getChannel(c, relayInfo, retryParam)
channel, selectedGroup, channelErr = getChannel(c, relayInfo, retryParam)
if channelErr != nil {
logger.LogError(c, channelErr.Error())
taskErr = service.TaskErrorWrapperLocal(channelErr.Err, "get_channel_failed", http.StatusInternalServerError)
@@ -522,6 +692,11 @@ func RelayTask(c *gin.Context) {
}
}

if bindErr := persistTaskVideoBindingIfNeeded(c.GetInt("id"), selectedGroup, channel); bindErr != nil {
taskErr = service.TaskErrorWrapperLocal(bindErr, "bind_task_video_channel_failed", http.StatusServiceUnavailable)
break
}

addUsedChannel(c, channel.Id)
bodyStorage, bodyErr := common.GetBodyStorage(c)
if bodyErr != nil {
@@ -569,14 +744,8 @@ func RelayTask(c *gin.Context) {
task.PrivateData.BillingSource = relayInfo.BillingSource
task.PrivateData.SubscriptionId = relayInfo.SubscriptionId
task.PrivateData.TokenId = relayInfo.TokenId
task.PrivateData.BillingContext = &model.TaskBillingContext{
ModelPrice: relayInfo.PriceData.ModelPrice,
GroupRatio: relayInfo.PriceData.GroupRatioInfo.GroupRatio,
ModelRatio: relayInfo.PriceData.ModelRatio,
OtherRatios: relayInfo.PriceData.OtherRatios,
OriginModelName: relayInfo.OriginModelName,
PerCallBilling: common.StringsContains(constant.TaskPricePatches, relayInfo.OriginModelName),
}
task.PrivateData.BillingContext = buildTaskBillingContext(relayInfo)
task.PrivateData.UpstreamRequest = relay.BuildUpstreamRequestSnapshotForTask(result.UpstreamReqJSON)
task.Quota = result.Quota
task.Data = result.TaskData
task.Action = relayInfo.Action
@@ -595,9 +764,34 @@ func respondTaskError(c *gin.Context, taskErr *dto.TaskError) {
if taskErr.StatusCode == http.StatusTooManyRequests {
taskErr.Message = "当前分组上游负载已饱和,请稍后再试"
}
if isKlingAipingNativePath(c.Request.URL.Path) {
normalizeKlingAipingTaskError(taskErr)
}
c.JSON(taskErr.StatusCode, taskErr)
}

func isKlingAipingNativePath(path string) bool {
return pathMatchesAnyKlingAipingNativePrefix(path,
"/v1/videos/text2video",
"/v1/videos/image2video",
"/v1/videos/motion-control",
"/v1/videos/omni-video",
"/v1/videos/multi-image2video",
"/v1/videos/video-extend",
"/v1/general/advanced-custom-elements",
"/v1/general/custom-voices",
)
}

func pathMatchesAnyKlingAipingNativePrefix(path string, prefixes ...string) bool {
for _, prefix := range prefixes {
if path == prefix || strings.HasPrefix(path, prefix+"/") {
return true
}
}
return false
}

func shouldRetryTaskRelay(c *gin.Context, channelId int, taskErr *dto.TaskError, retryTimes int) bool {
if taskErr == nil {
return false
@@ -608,9 +802,6 @@ func shouldRetryTaskRelay(c *gin.Context, channelId int, taskErr *dto.TaskError,
if retryTimes <= 0 {
return false
}
if _, ok := c.Get("specific_channel_id"); ok {
return false
}
if taskErr.StatusCode == http.StatusTooManyRequests {
return true
}


+ 152
- 0
controller/relay_test.go View File

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

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

"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
)

// ---------------------------------------------------------------------------
// shouldRetry — 渠道选择删除后的验证测试
// ---------------------------------------------------------------------------

func newTestContext(t *testing.T) *gin.Context {
t.Helper()
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request, _ = http.NewRequest("POST", "/v1/chat/completions", nil)
return c
}

func TestShouldRetry_NilError_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
assert.False(t, shouldRetry(c, nil, 3))
}

func TestShouldRetry_5xxErrorWithRetryLeft_ReturnsTrue(t *testing.T) {
// 渠道选择删除后,所有请求都能正常重试
// 之前设置了 specific_channel_id 时 shouldRetry 会返回 false
// 现在即使不设置任何 channel 相关 context,5xx 也应该重试
c := newTestContext(t)
apiErr := types.NewErrorWithStatusCode(
errors.New("upstream error"),
types.ErrorCodeBadResponseStatusCode,
http.StatusInternalServerError,
)
assert.True(t, shouldRetry(c, apiErr, 3))
}

func TestShouldRetry_5xxErrorZeroRetryLeft_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
apiErr := types.NewErrorWithStatusCode(
errors.New("upstream error"),
types.ErrorCodeBadResponseStatusCode,
http.StatusInternalServerError,
)
assert.False(t, shouldRetry(c, apiErr, 0))
}

func TestShouldRetry_SkipRetryError_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
apiErr := types.NewError(
errors.New("skip retry"),
types.ErrorCodeBadResponseStatusCode,
types.ErrOptionWithSkipRetry(),
)
assert.False(t, shouldRetry(c, apiErr, 3))
}

func TestShouldRetry_ChannelError_ReturnsTrue(t *testing.T) {
// channel: 前缀的错误码会触发立即重试
c := newTestContext(t)
apiErr := types.NewError(
errors.New("channel error"),
"channel:test",
)
// 即使 retryTimes=0,channel error 也返回 true
assert.True(t, shouldRetry(c, apiErr, 0))
}

func TestShouldRetry_2xxStatus_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
apiErr := types.NewErrorWithStatusCode(
errors.New("ok"),
types.ErrorCodeBadResponseStatusCode,
http.StatusOK,
)
assert.False(t, shouldRetry(c, apiErr, 3))
}

func TestShouldRetry_429Status_ReturnsTrue(t *testing.T) {
c := newTestContext(t)
apiErr := types.NewErrorWithStatusCode(
errors.New("rate limited"),
types.ErrorCodeBadResponseStatusCode,
http.StatusTooManyRequests,
)
assert.True(t, shouldRetry(c, apiErr, 3))
}

// ---------------------------------------------------------------------------
// shouldRetryTaskRelay — 渠道选择删除后的验证测试
// ---------------------------------------------------------------------------

func TestShouldRetryTaskRelay_NilError_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
assert.False(t, shouldRetryTaskRelay(c, 1, nil, 3))
}

func TestShouldRetryTaskRelay_5xxWithRetryLeft_ReturnsTrue(t *testing.T) {
c := newTestContext(t)
taskErr := &dto.TaskError{
Error: errors.New("server error"),
StatusCode: http.StatusInternalServerError,
}
assert.True(t, shouldRetryTaskRelay(c, 1, taskErr, 3))
}

func TestShouldRetryTaskRelay_5xxZeroRetryLeft_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
taskErr := &dto.TaskError{
Error: errors.New("server error"),
StatusCode: http.StatusInternalServerError,
}
assert.False(t, shouldRetryTaskRelay(c, 1, taskErr, 0))
}

func TestShouldRetryTaskRelay_429_ReturnsTrue(t *testing.T) {
c := newTestContext(t)
taskErr := &dto.TaskError{
Error: errors.New("rate limited"),
StatusCode: http.StatusTooManyRequests,
}
assert.True(t, shouldRetryTaskRelay(c, 1, taskErr, 1))
}

func TestShouldRetryTaskRelay_400_ReturnsFalse(t *testing.T) {
c := newTestContext(t)
taskErr := &dto.TaskError{
Error: errors.New("bad request"),
StatusCode: http.StatusBadRequest,
}
assert.False(t, shouldRetryTaskRelay(c, 1, taskErr, 3))
}

func TestShouldRetryTaskRelay_LocalError_ReturnsFalse(t *testing.T) {
// LocalError + 非 5xx 状态码(如 403)时,LocalError 检查生效返回 false
// 注意:5xx 分支在 LocalError 检查之前,所以 localError+5xx 仍会重试
c := newTestContext(t)
taskErr := &dto.TaskError{
Error: errors.New("local error"),
StatusCode: http.StatusForbidden,
LocalError: true,
}
assert.False(t, shouldRetryTaskRelay(c, 1, taskErr, 3))
}

+ 55
- 0
controller/reorder.go View File

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

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

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

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

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

+ 42
- 20
controller/topup.go View File

@@ -72,25 +72,25 @@ 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 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)
}
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,
@@ -164,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)


+ 10
- 63
controller/topup_alipay.go View File

@@ -3,19 +3,16 @@ package controller
import (
"context"
"encoding/base64"
"encoding/pem"
"fmt"
"log"
"net/http"
"strconv"
"strings"
"sync"
"time"

"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"
@@ -64,8 +61,11 @@ func getAlipayClient() (*alipay.Client, error) {
SetNotifyUrl(setting.AlipayNotifyURL)

// 设置支付宝公钥(用于回调验签)
pubKeyPEM := wrapAlipayPublicKey(setting.AlipayPublicKey)
client.AutoVerifySign([]byte(pubKeyPEM))
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
@@ -107,7 +107,7 @@ func RequestAlipayPayAmount(c *gin.Context) {
return
}

minTopup := getAlipayMinTopup()
minTopup := calcMinTopup(setting.AlipayMinTopUp)
if req.Amount < minTopup {
c.JSON(200, gin.H{"message": "error", "data": fmt.Sprintf("充值数量不能小于 %d", minTopup)})
return
@@ -120,7 +120,7 @@ func RequestAlipayPayAmount(c *gin.Context) {
return
}

payMoney := getAlipayPayMoney(float64(req.Amount), group)
payMoney := calcPayMoney(float64(req.Amount), group, setting.AlipayUnitPrice)
if payMoney <= 0.01 {
c.JSON(200, gin.H{"message": "error", "data": "充值金额过低"})
return
@@ -137,7 +137,7 @@ func RequestAlipayPay(c *gin.Context) {
return
}

minTopup := getAlipayMinTopup()
minTopup := calcMinTopup(setting.AlipayMinTopUp)
if req.Amount < minTopup {
c.JSON(200, gin.H{"message": "error", "data": fmt.Sprintf("充值数量不能小于 %d", minTopup)})
return
@@ -175,8 +175,8 @@ func RequestAlipayPay(c *gin.Context) {

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

payMoney := getAlipayPayMoney(float64(req.Amount), group)
totalAmount := strconv.FormatFloat(payMoney, 'f', 2, 64) // 支付宝金额单位是元
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 {
@@ -222,7 +222,6 @@ func AlipayPayStatus(c *gin.Context) {
return
}

// 验证订单属于当前用户
userId := c.GetInt("id")
if topUp.UserId != userId {
c.JSON(200, gin.H{"message": "error", "data": "订单不存在"})
@@ -261,7 +260,6 @@ func AlipayPayWebhook(c *gin.Context) {

tradeStatus := notifyReq.Get("trade_status")
if tradeStatus != "TRADE_SUCCESS" {
log.Printf("支付宝回调非成功状态: %s", tradeStatus)
c.String(http.StatusOK, "success")
return
}
@@ -280,54 +278,3 @@ func AlipayPayWebhook(c *gin.Context) {
log.Printf("支付宝充值成功: %s", tradeNo)
c.String(http.StatusOK, "success")
}

// getAlipayPayMoney 计算支付宝应付金额(元)
func getAlipayPayMoney(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.AlipayUnitPrice * topupGroupRatio * discount
return payMoney
}

// getAlipayMinTopup 获取支付宝最低充值数量
func getAlipayMinTopup() int64 {
minTopup := setting.AlipayMinTopUp
if operation_setting.GetQuotaDisplayType() == operation_setting.QuotaDisplayTypeTokens {
minTopup = minTopup * int(common.QuotaPerUnit)
}
return int64(minTopup)
}

// wrapAlipayPublicKey 将支付宝公钥包装为 PEM 格式
// 输入可能是:原始 Base64 字符串 或 已有 PEM 格式
func wrapAlipayPublicKey(pubKey string) string {
if strings.Contains(pubKey, "-----BEGIN") {
return pubKey
}
// 去除空白字符
cleaned := strings.ReplaceAll(pubKey, "\n", "")
cleaned = strings.ReplaceAll(cleaned, "\r", "")
cleaned = strings.TrimSpace(cleaned)
// Base64 解码为 DER 字节
derBytes, err := base64.StdEncoding.DecodeString(cleaned)
if err != nil {
// 如果解码失败,原样返回让上层报错
return pubKey
}
return string(pem.EncodeToMemory(&pem.Block{
Type: "PUBLIC KEY",
Bytes: derBytes,
}))
}

+ 2
- 24
controller/topup_wechat.go View File

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

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

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

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

+ 28
- 1
controller/user.go View File

@@ -14,6 +14,7 @@ import (
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/i18n"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/middleware"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/service/region_sync"
@@ -270,6 +271,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 +525,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": "",
@@ -557,6 +562,16 @@ func UpdateUser(c *gin.Context) {
common.ApiError(c, err)
return
}
if originUser.IsSyncedUser() {
// 对同步用户而言,界面展示和用户感知的“余额”是 synced_quota,
// 不能用本地 quota 列做比较,否则在 quota / synced_quota 漂移后会误判。
if originUser.SyncedQuota != updatedUser.Quota {
common.ApiErrorI18n(c, i18n.MsgSyncedUserQuotaCannotModify)
return
}
// 提交其他字段编辑时,保留本地 quota 原值,避免把展示用 synced_quota 回写到 quota 列。
updatedUser.Quota = originUser.Quota
}
myRole := c.GetInt("role")
if myRole <= originUser.Role && myRole != common.RoleRootUser {
common.ApiErrorI18n(c, i18n.MsgUserNoPermissionHigherLevel)
@@ -574,6 +589,7 @@ func UpdateUser(c *gin.Context) {
common.ApiError(c, err)
return
}
middleware.SetCaptureEnabled(int64(updatedUser.Id), updatedUser.CaptureRelay)
if originUser.Quota != updatedUser.Quota {
model.RecordLog(originUser.Id, model.LogTypeManage, fmt.Sprintf("管理员将用户额度从 %s修改为 %s", logger.LogQuota(originUser.Quota), logger.LogQuota(updatedUser.Quota)))
}
@@ -827,12 +843,19 @@ func CreateUser(c *gin.Context) {
Password: user.Password,
DisplayName: user.DisplayName,
Role: user.Role, // 保持管理员设置的角色
Email: user.Email,
Quota: user.Quota,
Group: user.Group,
AffCode: user.AffCode,
}
if err := cleanUser.Insert(0); err != nil {
common.ApiError(c, err)
return
}

// 同步用户到海外节点(异步执行,不阻塞创建流程)
region_sync.PushUserCreateToSlave(&cleanUser)

c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
@@ -1019,6 +1042,10 @@ func TopUp(c *gin.Context) {
}
quota, err := model.Redeem(req.Key, id)
if err != nil {
if errors.Is(err, model.ErrSyncedUserRedeemDenied) {
common.ApiErrorI18n(c, i18n.MsgRedemptionSyncedUserDenied)
return
}
if errors.Is(err, model.ErrRedeemFailed) {
common.ApiErrorI18n(c, i18n.MsgRedeemFailed)
return


+ 376
- 0
controller/user_migration.go View File

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

import (
"errors"
"net/http"
"strconv"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
user_migration "github.com/QuantumNous/new-api/service/user_migration"
"github.com/QuantumNous/new-api/setting/system_setting"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)

type migrationVerifier interface {
VerifyItem(item *model.UserMigrationItem) (*user_migration.VerifyResult, error)
}

type migrationCandidateClient interface {
ListUsers(page, pageSize int, keyword string) (*user_migration.ListRemoteUsersResponse, error)
}

var newMigrationVerifier = func() migrationVerifier {
return newMigrationService()
}

var newMigrationCandidateClient = func() migrationCandidateClient {
settings := system_setting.GetRegionSyncSettings()
endpoint := settings.MasterEndpoint
if settings.IsMaster && len(settings.SlaveEndpoints) > 0 {
endpoint = settings.SlaveEndpoints[0]
}
return user_migration.NewClient(endpoint, settings.SyncApiKey)
}

func CreateUserMigrationBatch(c *gin.Context) {
var req struct {
Name string `json:"name" binding:"required"`
SourceRegion string `json:"source_region" binding:"required"`
TargetRegion string `json:"target_region" binding:"required"`
SelectionMode string `json:"selection_mode"`
SourceUserIDs []int `json:"source_user_ids"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": err.Error()})
return
}

operatorID, _ := c.Get("id")
if req.SelectionMode == "" {
req.SelectionMode = user_migration.SelectionModeAll
}
if req.SelectionMode != user_migration.SelectionModeAll && req.SelectionMode != user_migration.SelectionModeExplicitIDs {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid selection_mode"})
return
}

payload := user_migration.SelectionPayload{}
if req.SelectionMode == user_migration.SelectionModeExplicitIDs {
if len(req.SourceUserIDs) == 0 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "source_user_ids is required"})
return
}
seen := make(map[int]struct{}, len(req.SourceUserIDs))
payload.SourceUserIDs = make([]int, 0, len(req.SourceUserIDs))
for _, id := range req.SourceUserIDs {
if id <= 0 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "invalid source_user_ids"})
return
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
payload.SourceUserIDs = append(payload.SourceUserIDs, id)
}
if len(payload.SourceUserIDs) > 1000 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "too many source_user_ids"})
return
}

conflicts, err := model.FindActiveMigrationUserConflicts(payload.SourceUserIDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
if len(conflicts) > 0 {
c.JSON(http.StatusConflict, gin.H{
"success": false,
"error": "some selected users already exist in active migration batches",
"conflicts": conflicts,
})
return
}
}
payloadJSON, _ := common.Marshal(payload)
batch := &model.UserMigrationBatch{
Name: req.Name,
SourceRegion: req.SourceRegion,
TargetRegion: req.TargetRegion,
SelectionMode: req.SelectionMode,
SelectionPayload: string(payloadJSON),
RequestedUserCount: len(payload.SourceUserIDs),
OperatorId: operatorID.(int),
}
if err := model.CreateUserMigrationBatch(batch); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "data": batch})
}

func ListMigrationCandidateUsers(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "0"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
keyword := c.Query("keyword")
if pageSize <= 0 || pageSize > 200 {
pageSize = 50
}

resp, err := newMigrationCandidateClient().ListUsers(page, pageSize, keyword)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"data": resp.Data,
"total": resp.Total,
})
}

func ListUserMigrationBatches(c *gin.Context) {
var batches []model.UserMigrationBatch
if err := model.DB.Order("id desc").Find(&batches).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "data": batches})
}

func GetUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))

var batch model.UserMigrationBatch
if err := model.DB.First(&batch, batchID).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "batch not found"})
return
}

page, _ := strconv.Atoi(c.DefaultQuery("page", "0"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
if pageSize <= 0 {
pageSize = 50
}
if pageSize > 200 {
pageSize = 200
}

var items []model.UserMigrationItem
var total int64
model.DB.Model(&model.UserMigrationItem{}).Where("batch_id = ?", batchID).Count(&total)
model.DB.Where("batch_id = ?", batchID).Order("id asc").Offset(page * pageSize).Limit(pageSize).Find(&items)

c.JSON(http.StatusOK, gin.H{
"success": true,
"data": batch,
"items": items,
"total": total,
})
}

func ScanUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))
locked, err := model.TryLockBatchForScan(batchID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
if !locked {
c.JSON(http.StatusConflict, gin.H{"success": false, "error": "batch is already being scanned or not in a scannable state"})
return
}

svc := newMigrationService()
go func() {
if err := svc.ScanBatch(batchID); err != nil {
model.DB.Model(&model.UserMigrationBatch{}).Where("id = ?", batchID).
Update("status", model.UserMigrationBatchStatusFailed)
}
}()

c.JSON(http.StatusAccepted, gin.H{"success": true, "message": "scan started"})
}

func ResolveUserMigrationItem(c *gin.Context) {
itemID, _ := strconv.Atoi(c.Param("id"))

var item model.UserMigrationItem
if err := model.DB.Select("id", "batch_id", "conflict_flags").First(&item, itemID).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "item not found"})
return
}

var req struct {
ResolutionStrategy string `json:"resolution_strategy" binding:"required"`
TargetUserId int `json:"target_user_id"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": err.Error()})
return
}

validStrategies := map[string]bool{
model.UserMigrationStrategyCreateNew: true,
model.UserMigrationStrategyMergeExisting: true,
model.UserMigrationStrategySkip: true,
}
if !validStrategies[req.ResolutionStrategy] {
c.JSON(http.StatusBadRequest, gin.H{
"success": false,
"error": "invalid resolution_strategy: must be create_new, merge_into_existing, or skip",
})
return
}
if req.ResolutionStrategy == model.UserMigrationStrategyMergeExisting && req.TargetUserId == 0 {
c.JSON(http.StatusBadRequest, gin.H{
"success": false,
"error": "merge_into_existing requires target_user_id",
})
return
}
if req.ResolutionStrategy == model.UserMigrationStrategyMergeExisting &&
hasMigrationConflictFlag(item.ConflictFlags, "cn_already_synced_to_ov") {
c.JSON(http.StatusConflict, gin.H{
"success": false,
"error": "merge_into_existing is not allowed when target CN user already has a synced OV copy",
})
return
}

updates := map[string]any{
"resolution_strategy": req.ResolutionStrategy,
"status": model.UserMigrationItemStatusReady,
}
if req.TargetUserId != 0 {
updates["target_user_id"] = req.TargetUserId
}
if err := model.DB.Model(&model.UserMigrationItem{}).Where("id = ?", itemID).Updates(updates).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
if err := model.RefreshUserMigrationBatchStats(item.BatchId); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}

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

func hasMigrationConflictFlag(flagsJSON string, want string) bool {
if flagsJSON == "" || want == "" {
return false
}
var flags []string
if err := common.UnmarshalJsonStr(flagsJSON, &flags); err != nil {
return false
}
for _, flag := range flags {
if flag == want {
return true
}
}
return false
}

func ExecuteUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))
locked, err := model.TryLockBatchForExecution(batchID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
if !locked {
c.JSON(http.StatusConflict, gin.H{"success": false, "error": "batch is already running or not in an executable state"})
return
}

svc := newMigrationService()
go func() {
_ = svc.ExecuteBatch(batchID)
}()

c.JSON(http.StatusAccepted, gin.H{"success": true, "message": "execution started"})
}

func RetryUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))

model.DB.Model(&model.UserMigrationItem{}).
Where("batch_id = ? AND status = ?", batchID, model.UserMigrationItemStatusFailed).
Updates(map[string]any{"status": model.UserMigrationItemStatusReady, "error_message": ""})

locked, err := model.TryLockBatchForExecution(batchID)
if err != nil || !locked {
c.JSON(http.StatusConflict, gin.H{"success": false, "error": "batch cannot be retried now"})
return
}

svc := newMigrationService()
go func() {
_ = svc.ExecuteBatch(batchID)
}()

c.JSON(http.StatusAccepted, gin.H{"success": true, "message": "retry started"})
}

func CancelUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))

cancelled, err := model.CancelUserMigrationBatch(batchID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "batch not found"})
return
}
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": err.Error()})
return
}
if !cancelled {
c.JSON(http.StatusConflict, gin.H{
"success": false,
"error": "batch cannot be cancelled in current status",
})
return
}

c.JSON(http.StatusOK, gin.H{"success": true, "message": "batch cancelled"})
}

func VerifyUserMigrationBatch(c *gin.Context) {
batchID, _ := strconv.Atoi(c.Param("id"))

var items []model.UserMigrationItem
model.DB.Where("batch_id = ? AND status = ?", batchID, model.UserMigrationItemStatusMigrated).Find(&items)

svc := newMigrationVerifier()
results := make([]map[string]any, 0, len(items))
for _, item := range items {
result, err := svc.VerifyItem(&item)
entry := map[string]any{
"source_user_id": item.SourceUserId,
"target_user_id": item.TargetUserId,
}
if err != nil {
entry["error"] = err.Error()
} else {
entry["result"] = result
}
results = append(results, entry)
}

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

func newMigrationService() *user_migration.Service {
settings := system_setting.GetRegionSyncSettings()
endpoint := settings.MasterEndpoint
if settings.IsMaster && len(settings.SlaveEndpoints) > 0 {
endpoint = settings.SlaveEndpoints[0]
}

client := user_migration.NewClient(endpoint, settings.SyncApiKey)
return user_migration.NewService(client)
}

+ 839
- 0
controller/user_migration_test.go View File

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

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

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

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

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

sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

require.NoError(t, db.AutoMigrate(
&model.User{},
&model.CustomOAuthProvider{},
&model.UserOAuthBinding{},
&model.UserMigrationBatch{},
&model.UserMigrationItem{},
&model.MigrationQuotaGrant{},
))

return db
}

func testRootAuthMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
c.Set("id", 1)
c.Set("role", 100)
c.Next()
}
}

type fakeMigrationCandidateClient struct {
resp *user_migration.ListRemoteUsersResponse
err error
gotKeyword string
gotPage int
gotPageSize int
}

func (f *fakeMigrationCandidateClient) ListUsers(page, pageSize int, keyword string) (*user_migration.ListRemoteUsersResponse, error) {
f.gotKeyword = keyword
f.gotPage = page
f.gotPageSize = pageSize
if f.err != nil {
return nil, f.err
}
return f.resp, nil
}

type fakeMigrationVerifier struct {
results map[int]*user_migration.VerifyResult
errs map[int]error
}

func (f *fakeMigrationVerifier) VerifyItem(item *model.UserMigrationItem) (*user_migration.VerifyResult, error) {
if err, ok := f.errs[item.SourceUserId]; ok {
return nil, err
}
if result, ok := f.results[item.SourceUserId]; ok {
return result, nil
}
return nil, errors.New("unexpected source user")
}

func TestCreateUserMigrationBatch_Success(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches", CreateUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

body := map[string]any{
"name": "wave-root-1",
"source_region": "overseas",
"target_region": "cn",
}
bodyBytes, err := json.Marshal(body)
require.NoError(t, err)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

func TestCreateUserMigrationBatch_ExplicitIDsSuccess(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches", CreateUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches", bytes.NewBufferString(`{
"name":"explicit-batch",
"source_region":"ov",
"target_region":"cn",
"selection_mode":"explicit_ids",
"source_user_ids":[1001,1002,1002]
}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var batch model.UserMigrationBatch
require.NoError(t, db.First(&batch, 1).Error)
require.Equal(t, "explicit_ids", batch.SelectionMode)
require.Equal(t, 2, batch.RequestedUserCount)
}

func TestCreateUserMigrationBatch_ExplicitIDsValidation(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches", CreateUserMigrationBatch)

cases := []string{
`{"name":"b1","source_region":"ov","target_region":"cn","selection_mode":"explicit_ids","source_user_ids":[]}`,
`{"name":"b2","source_region":"ov","target_region":"cn","selection_mode":"explicit_ids","source_user_ids":[0]}`,
`{"name":"b3","source_region":"ov","target_region":"cn","selection_mode":"invalid","source_user_ids":[1]}`,
}

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

for _, body := range cases {
req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches", bytes.NewBufferString(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusBadRequest, w.Code)
}
}

func TestCreateUserMigrationBatch_ExplicitIDs_RejectsActiveBatchOverlap(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches", CreateUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

existingBatch := &model.UserMigrationBatch{
Name: "existing-draft-batch",
SourceRegion: "ov",
TargetRegion: "cn",
SelectionMode: "explicit_ids",
SelectionPayload: `{"source_user_ids":[1002]}`,
Status: model.UserMigrationBatchStatusDraft,
OperatorId: 1,
}
require.NoError(t, db.Create(existingBatch).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches", bytes.NewBufferString(`{
"name":"explicit-batch",
"source_region":"ov",
"target_region":"cn",
"selection_mode":"explicit_ids",
"source_user_ids":[1001,1002]
}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusConflict, w.Code)
var resp struct {
Success bool `json:"success"`
Error string `json:"error"`
Conflicts []struct {
SourceUserID int `json:"source_user_id"`
BatchID int `json:"batch_id"`
BatchStatus string `json:"batch_status"`
} `json:"conflicts"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.False(t, resp.Success)
require.Equal(t, "some selected users already exist in active migration batches", resp.Error)
require.Len(t, resp.Conflicts, 1)
require.Equal(t, 1002, resp.Conflicts[0].SourceUserID)
require.Equal(t, existingBatch.Id, resp.Conflicts[0].BatchID)
require.Equal(t, model.UserMigrationBatchStatusDraft, resp.Conflicts[0].BatchStatus)
}

func TestListMigrationCandidateUsers_Success(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())

origFactory := newMigrationCandidateClient
defer func() {
newMigrationCandidateClient = origFactory
}()
newMigrationCandidateClient = func() migrationCandidateClient {
return &fakeMigrationCandidateClient{
resp: &user_migration.ListRemoteUsersResponse{
Success: true,
Data: []user_migration.RemoteUserSnapshot{
{Id: 10000001, Username: "ov-a", Email: "a@example.com", Quota: 100},
},
Total: 1,
},
}
}

router.GET("/api/user-migrations/candidate-users", ListMigrationCandidateUsers)
req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/candidate-users?page=0&page_size=50", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

func TestListMigrationCandidateUsers_UpstreamFailure(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())

origFactory := newMigrationCandidateClient
defer func() {
newMigrationCandidateClient = origFactory
}()
newMigrationCandidateClient = func() migrationCandidateClient {
return &fakeMigrationCandidateClient{err: errors.New("upstream unavailable")}
}

router.GET("/api/user-migrations/candidate-users", ListMigrationCandidateUsers)
req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/candidate-users?page=0&page_size=50", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusInternalServerError, w.Code)
}

func TestListMigrationCandidateUsers_ForwardsKeyword(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())

client := &fakeMigrationCandidateClient{
resp: &user_migration.ListRemoteUsersResponse{
Success: true,
Data: []user_migration.RemoteUserSnapshot{},
Total: 0,
},
}

origFactory := newMigrationCandidateClient
defer func() {
newMigrationCandidateClient = origFactory
}()
newMigrationCandidateClient = func() migrationCandidateClient {
return client
}

router.GET("/api/user-migrations/candidate-users", ListMigrationCandidateUsers)
req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/candidate-users?page=2&page_size=30&keyword=alice", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
require.Equal(t, 2, client.gotPage)
require.Equal(t, 30, client.gotPageSize)
require.Equal(t, "alice", client.gotKeyword)
}

func TestListUserMigrationBatches_Success(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.GET("/api/user-migrations/batches", ListUserMigrationBatches)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "batch-1",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 2,
Name: "batch-2",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/batches", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

func TestGetUserMigrationBatch_Success(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.GET("/api/user-migrations/batches/:id", GetUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "detail-batch",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
BatchId: 1,
SourceUserId: 1001,
SourceUsername: "ov-u1",
Status: model.UserMigrationItemStatusReady,
}).Error)

req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/batches/1?page=0&page_size=50", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusOK, w.Code)
var resp map[string]any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.True(t, resp["success"].(bool))
require.NotNil(t, resp["data"])
require.Len(t, resp["items"], 1)
require.EqualValues(t, 1, resp["total"])
}

func TestGetUserMigrationBatch_NotFound(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.GET("/api/user-migrations/batches/:id", GetUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

req := httptest.NewRequest(http.MethodGet, "/api/user-migrations/batches/999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusNotFound, w.Code)
}

func TestExecuteUserMigrationBatch_AlreadyRunning(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/execute", ExecuteUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

batch := &model.UserMigrationBatch{
Name: "b",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
Status: model.UserMigrationBatchStatusRunning,
}
require.NoError(t, db.Create(batch).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/execute", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusConflict, w.Code)
}

func TestExecuteUserMigrationBatch_InvalidStateRejected(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/execute", ExecuteUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

for _, status := range []string{model.UserMigrationBatchStatusDraft, model.UserMigrationBatchStatusScanned} {
require.NoError(t, db.Exec("DELETE FROM user_migration_batches").Error)
require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "invalid-exec",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
Status: status,
}).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/execute", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)
}
}

func TestResolveUserMigrationItem_InvalidStrategy(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/items/:id/resolve", ResolveUserMigrationItem)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationItem{
BatchId: 1,
SourceUserId: 10000001,
SourceUsername: "u",
Status: model.UserMigrationItemStatusConflict,
}).Error)

req := httptest.NewRequest(
http.MethodPost,
"/api/user-migrations/items/1/resolve",
bytes.NewBufferString(`{"resolution_strategy":"invalid_strategy"}`),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusBadRequest, w.Code)
}

func TestResolveUserMigrationItem_RefreshesBatchStatusToReady(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/items/:id/resolve", ResolveUserMigrationItem)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

batch := &model.UserMigrationBatch{
Name: "resolve-ready",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
Status: model.UserMigrationBatchStatusScanned,
}
require.NoError(t, model.CreateUserMigrationBatch(batch))
require.NoError(t, db.Create(&model.UserMigrationItem{
BatchId: batch.Id,
SourceUserId: 1001,
SourceUsername: "ov-u1",
Status: model.UserMigrationItemStatusConflict,
}).Error)

req := httptest.NewRequest(
http.MethodPost,
"/api/user-migrations/items/1/resolve",
bytes.NewBufferString(`{"resolution_strategy":"create_new"}`),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

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

var refreshed model.UserMigrationBatch
require.NoError(t, db.First(&refreshed, batch.Id).Error)
require.Equal(t, model.UserMigrationBatchStatusReady, refreshed.Status)
}

func TestResolveUserMigrationItem_RejectsMergeWhenCNAlreadyHasSyncedCopy(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/items/:id/resolve", ResolveUserMigrationItem)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 1,
BatchId: 1,
SourceUserId: 1001,
SourceUsername: "ov-u1",
Status: model.UserMigrationItemStatusConflict,
ConflictFlags: `["email","cn_already_synced_to_ov"]`,
TargetUserId: 24,
}).Error)

req := httptest.NewRequest(
http.MethodPost,
"/api/user-migrations/items/1/resolve",
bytes.NewBufferString(`{"resolution_strategy":"merge_into_existing","target_user_id":24}`),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

require.Equal(t, http.StatusConflict, w.Code)

var item model.UserMigrationItem
require.NoError(t, db.First(&item, 1).Error)
require.Equal(t, model.UserMigrationItemStatusConflict, item.Status)
require.Equal(t, 24, item.TargetUserId)
}

func TestRetryUserMigrationBatch_OnlyFailedItemsReset(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/retry", RetryUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "retry-batch",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
Status: model.UserMigrationBatchStatusRunning,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 1,
BatchId: 1,
SourceUserId: 1001,
SourceUsername: "failed-user",
Status: model.UserMigrationItemStatusFailed,
ErrorMessage: "boom",
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 2,
BatchId: 1,
SourceUserId: 1002,
SourceUsername: "migrated-user",
Status: model.UserMigrationItemStatusMigrated,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 3,
BatchId: 1,
SourceUserId: 1003,
SourceUsername: "skipped-user",
Status: model.UserMigrationItemStatusSkipped,
}).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/retry", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)

var failedItem, migratedItem, skippedItem model.UserMigrationItem
require.NoError(t, db.First(&failedItem, 1).Error)
require.NoError(t, db.First(&migratedItem, 2).Error)
require.NoError(t, db.First(&skippedItem, 3).Error)

require.Equal(t, model.UserMigrationItemStatusReady, failedItem.Status)
require.Empty(t, failedItem.ErrorMessage)
require.Equal(t, model.UserMigrationItemStatusMigrated, migratedItem.Status)
require.Equal(t, model.UserMigrationItemStatusSkipped, skippedItem.Status)
}

func TestCancelUserMigrationBatch_SoftDeleteSuccess(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/cancel", CancelUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "cancel-me",
SourceRegion: "ov",
TargetRegion: "cn",
SelectionMode: "explicit_ids",
SelectionPayload: `{"source_user_ids":[11,12]}`,
OperatorId: 1,
Status: model.UserMigrationBatchStatusScanned,
}).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/cancel", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)

var batch model.UserMigrationBatch
require.NoError(t, db.First(&batch, 1).Error)
require.Equal(t, model.UserMigrationBatchStatusCancelled, batch.Status)
}

func TestCancelUserMigrationBatch_InvalidStateRejected(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/cancel", CancelUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

for _, status := range []string{
model.UserMigrationBatchStatusScanning,
model.UserMigrationBatchStatusRunning,
model.UserMigrationBatchStatusCompleted,
} {
require.NoError(t, db.Exec("DELETE FROM user_migration_batches").Error)
require.NoError(t, db.Create(&model.UserMigrationBatch{
Id: 1,
Name: "cannot-cancel",
SourceRegion: "ov",
TargetRegion: "cn",
OperatorId: 1,
Status: status,
}).Error)

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/cancel", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusConflict, w.Code)
}
}

func TestCreateUserMigrationBatch_ExplicitIDs_AllowsReuseAfterCancel(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches", CreateUserMigrationBatch)
router.POST("/api/user-migrations/batches/:id/cancel", CancelUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
model.DB = db
defer func() {
model.DB = orig
}()

existingBatch := &model.UserMigrationBatch{
Name: "existing-scanned-batch",
SourceRegion: "ov",
TargetRegion: "cn",
SelectionMode: "explicit_ids",
SelectionPayload: `{"source_user_ids":[11,12]}`,
Status: model.UserMigrationBatchStatusScanned,
OperatorId: 1,
}
require.NoError(t, db.Create(existingBatch).Error)

cancelReq := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/cancel", nil)
cancelW := httptest.NewRecorder()
router.ServeHTTP(cancelW, cancelReq)
require.Equal(t, http.StatusOK, cancelW.Code)

createReq := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches", bytes.NewBufferString(`{
"name":"recreated-batch",
"source_region":"ov",
"target_region":"cn",
"selection_mode":"explicit_ids",
"source_user_ids":[11,12]
}`))
createReq.Header.Set("Content-Type", "application/json")
createW := httptest.NewRecorder()
router.ServeHTTP(createW, createReq)
require.Equal(t, http.StatusOK, createW.Code)
}

func TestVerifyUserMigrationBatch_AggregatesResult(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.POST("/api/user-migrations/batches/:id/verify", VerifyUserMigrationBatch)

db := setupUserMigrationControllerDB(t)
orig := model.DB
origFactory := newMigrationVerifier
model.DB = db
defer func() {
model.DB = orig
newMigrationVerifier = origFactory
}()

require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 1,
BatchId: 1,
SourceUserId: 1001,
TargetUserId: 1,
Status: model.UserMigrationItemStatusMigrated,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 2,
BatchId: 1,
SourceUserId: 1002,
TargetUserId: 2,
Status: model.UserMigrationItemStatusMigrated,
}).Error)
require.NoError(t, db.Create(&model.UserMigrationItem{
Id: 3,
BatchId: 1,
SourceUserId: 1003,
TargetUserId: 3,
Status: model.UserMigrationItemStatusMigrated,
}).Error)

newMigrationVerifier = func() migrationVerifier {
return &fakeMigrationVerifier{
results: map[int]*user_migration.VerifyResult{
1001: {
TargetUserExists: true,
RemoteConverted: true,
QuotaMatched: true,
},
1002: {
TargetUserExists: false,
RemoteConverted: true,
QuotaMatched: false,
},
},
errs: map[int]error{
1003: errors.New("remote user not found"),
},
}
}

req := httptest.NewRequest(http.MethodPost, "/api/user-migrations/batches/1/verify", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)

var resp struct {
Success bool `json:"success"`
Data []map[string]any `json:"data"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.True(t, resp.Success)
require.Len(t, resp.Data, 3)

bySource := make(map[int]map[string]any, len(resp.Data))
for _, entry := range resp.Data {
bySource[int(entry["source_user_id"].(float64))] = entry
}

require.NotNil(t, bySource[1001]["result"])
require.NotNil(t, bySource[1002]["result"])
require.Equal(t, "remote user not found", bySource[1003]["error"])
}

+ 74
- 0
controller/user_rate_limit.go View File

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

import (
"strconv"

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

// GetUserRateLimits GET /api/user/:id/rate-limits
// 查询指定用户的所有模型 RPM 限制配置(管理员权限)
func GetUserRateLimits(c *gin.Context) {
userId, err := strconv.Atoi(c.Param("id"))
if err != nil {
common.ApiErrorMsg(c, "invalid user id")
return
}

list, err := model.GetUserModelRateLimits(userId)
if err != nil {
common.ApiError(c, err)
return
}

// 保证返回空列表而非 null
if list == nil {
list = []model.UserModelRateLimit{}
}

c.JSON(200, gin.H{
"success": true,
"message": "",
"data": list,
})
}

// SetUserRateLimits PUT /api/user/:id/rate-limits
// 覆盖式写入指定用户的所有模型 RPM 限制配置(管理员权限)
func SetUserRateLimits(c *gin.Context) {
userId, err := strconv.Atoi(c.Param("id"))
if err != nil {
common.ApiErrorMsg(c, "invalid user id")
return
}

var items []model.UserModelRateLimit
if err := c.ShouldBindJSON(&items); err != nil {
common.ApiErrorMsg(c, err.Error())
return
}

for _, item := range items {
if item.Rpm < 0 {
common.ApiErrorMsg(c, "rpm must be >= 0")
return
}
}

if err := model.SetUserModelRateLimits(userId, items); err != nil {
common.ApiError(c, err)
return
}

// 删除 Redis 缓存
if common.RedisEnabled {
_ = common.RedisDel("user_model_rate_limit:" + strconv.Itoa(userId))
}

c.JSON(200, gin.H{
"success": true,
"message": "",
})
}

+ 152
- 0
controller/user_rate_limit_test.go View File

@@ -0,0 +1,152 @@
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 setupUserRateLimitControllerDB(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 := model.DB
model.DB = db
common.RedisEnabled = false

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

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

func setupUserRateLimitRouter() *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
g := r.Group("/api/user")
{
g.GET("/:id/rate-limits", GetUserRateLimits)
g.PUT("/:id/rate-limits", SetUserRateLimits)
}
return r
}

func TestGetUserRateLimits_Empty(t *testing.T) {
setupUserRateLimitControllerDB(t)
router := setupUserRateLimitRouter()

w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/api/user/1/rate-limits", 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))
// data 应为空列表(非 nil)
data, ok := resp["data"]
require.True(t, ok, "response should have data key")
assert.Empty(t, data)
}

func TestSetUserRateLimits_OK(t *testing.T) {
setupUserRateLimitControllerDB(t)
router := setupUserRateLimitRouter()

items := []model.UserModelRateLimit{
{Model: "gpt-4", Rpm: 60},
}
body, _ := json.Marshal(items)

w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodPut, "/api/user/1/rate-limits", 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))

// 验证数据确实写入
list, err := model.GetUserModelRateLimits(1)
require.NoError(t, err)
require.Len(t, list, 1)
assert.Equal(t, "gpt-4", list[0].Model)
assert.Equal(t, 60, list[0].Rpm)
}

func TestSetUserRateLimits_NegativeRpm(t *testing.T) {
setupUserRateLimitControllerDB(t)
router := setupUserRateLimitRouter()

items := []model.UserModelRateLimit{
{Model: "gpt-4", Rpm: -1},
}
body, _ := json.Marshal(items)

w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodPut, "/api/user/1/rate-limits", 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.False(t, resp["success"].(bool))
}

func TestGetUserRateLimits_InvalidId(t *testing.T) {
setupUserRateLimitControllerDB(t)
router := setupUserRateLimitRouter()

w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodGet, "/api/user/notanid/rate-limits", 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.False(t, resp["success"].(bool))
}

func TestSetUserRateLimits_EmptySlice(t *testing.T) {
setupUserRateLimitControllerDB(t)
router := setupUserRateLimitRouter()

// 先写一条
_ = model.SetUserModelRateLimits(2, []model.UserModelRateLimit{
{UserId: 2, Model: "claude-3-opus", Rpm: 10},
})

// 用空 slice 清空
body, _ := json.Marshal([]model.UserModelRateLimit{})
w := httptest.NewRecorder()
req, _ := http.NewRequest(http.MethodPut, "/api/user/2/rate-limits", 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))

list, err := model.GetUserModelRateLimits(2)
require.NoError(t, err)
assert.Empty(t, list)
}

+ 146
- 0
controller/user_sync_guard_test.go View File

@@ -0,0 +1,146 @@
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/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

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

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

sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

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

func TestUpdateUser_SyncedUserAllowsNonQuotaEditsWithoutOverwritingQuota(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.PUT("/api/user/", UpdateUser)

db := setupUserControllerDB(t)
orig := model.DB
origRedisEnabled := common.RedisEnabled
model.DB = db
defer func() {
model.DB = orig
common.RedisEnabled = origRedisEnabled
}()
common.RedisEnabled = false

require.NoError(t, db.Create(&model.User{
Id: 1001,
Username: "synced-user",
Password: "hashed-password",
DisplayName: "Old Name",
Role: common.RoleCommonUser,
Status: common.UserStatusEnabled,
Group: "default",
Quota: 0,
SyncedQuota: 100000,
Source: common.UserSourceSynced,
AffCode: "SYNC1",
}).Error)

body := map[string]any{
"id": 1001,
"username": "synced-user",
"display_name": "New Name",
"password": "",
"group": "default",
"quota": 100000,
}
bodyBytes, err := json.Marshal(body)
require.NoError(t, err)

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

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

var saved model.User
require.NoError(t, db.First(&saved, 1001).Error)
require.Equal(t, "New Name", saved.DisplayName)
require.Equal(t, 0, saved.Quota)
require.Equal(t, 100000, saved.SyncedQuota)
}

func TestUpdateUser_SyncedUserQuotaChangeRejected(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(testRootAuthMiddleware())
router.PUT("/api/user/", UpdateUser)

db := setupUserControllerDB(t)
orig := model.DB
origRedisEnabled := common.RedisEnabled
model.DB = db
defer func() {
model.DB = orig
common.RedisEnabled = origRedisEnabled
}()
common.RedisEnabled = false

require.NoError(t, db.Create(&model.User{
Id: 1002,
Username: "synced-user-2",
Password: "hashed-password",
DisplayName: "Synced User 2",
Role: common.RoleCommonUser,
Status: common.UserStatusEnabled,
Group: "default",
Quota: 0,
SyncedQuota: 100000,
Source: common.UserSourceSynced,
AffCode: "SYNC2",
}).Error)

body := map[string]any{
"id": 1002,
"username": "synced-user-2",
"display_name": "Synced User 2",
"password": "",
"group": "default",
"quota": 100001,
}
bodyBytes, err := json.Marshal(body)
require.NoError(t, err)

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

require.Equal(t, http.StatusOK, w.Code)
var resp map[string]any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.False(t, resp["success"].(bool))
require.Equal(t, "user.synced_user_quota_cannot_modify", resp["message"])

var saved model.User
require.NoError(t, db.First(&saved, 1002).Error)
require.Equal(t, 0, saved.Quota)
require.Equal(t, 100000, saved.SyncedQuota)
}

+ 106
- 0
controller/user_topup_test.go View File

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

import (
"bytes"
"net/http"