| Author | SHA1 | Message | Date |
|---|---|---|---|
|
|
4533c24fa1 | merge: dynamic video asset channel selection | 3 days ago |
|
|
f6321b3910 |
feat(channel): support dynamic video asset channel selection
Co-Authored-By: Codex <noreply@anthropic.com> |
3 days ago |
|
|
3fbf48313a |
feat(pricing): support usage billing and endpoint fixes
Add KlingAiping usage-token settlement from upstream final_unit_deduction and duration fallback handling. Normalize model endpoint config persistence and convert model pricing values between stored USD and display currency in the editor. Co-Authored-By: Codex <noreply@anthropic.com> |
3 days ago |
|
|
dd14839996 | merge: tianyiyun seedance channel | 3 days ago |
|
|
678add297f |
feat(channel): add tianyiyun seedance channel
Co-Authored-By: Codex <noreply@anthropic.com> |
3 days ago |
|
|
f100f9e409 | docs: design tianyiyun seedance channel | 3 days ago |
|
|
ffb3991654 | Merge branch 'feat/kling-aiping' | 6 days ago |
|
|
49ca8a6ae7 |
fix(kling): 修复原生路由和展示价格保存
- 避免仅编辑展示价格时写入默认真实计费规则 - 将 video-extend 纳入 KlingAiping 原生渠道约束 - 补充相关保存决策、渠道约束和计费展示测试 Co-Authored-By: Codex <noreply@anthropic.com> |
6 days ago |
|
|
6f805a6177 |
refactor: update aiping doubao upstream API paths to multimodal/sd endpoints
- Asset: /api/v1/volcengine/asset → /api/v1/multimodal/sd/assets
- Create video task: /api/v1/videos → /api/v1/multimodal/sd/videos/contents/generations/tasks
- Query video task: /api/v1/videos/{id} → /api/v1/multimodal/sd/videos/contents/generations/tasks/{id}
Co-Authored-By: Claude <noreply@anthropic.com>
|
1 week ago |
|
|
56400b6c61 |
feat: KlingAiping 渠道适配 + 展示价格管理后台
- 新增 KlingAiping 渠道 (ChannelTypeKlingAiping = 59) - 新增 KlingAiping 原生路由 (text2video/image2video/omni-video 等) - 新增 KlingAiping adaptor (relay/channel/task/kling/aiping/) - 前端模型定价页面新增「真实计费 / 展示价格」Tab 切换 - 新增 DisplayPricing 组件,支持模型展示价格的增删改 - 新增 displayPricingConfig 工具函数及测试 - 修复 KlingAiping 路由路径匹配与 task 错误响应 Co-Authored-By: Claude <noreply@anthropic.com> |
1 week ago |
|
|
a61ca078fc | fix: localize display pricing units | 1 week ago |
|
|
fee9504107 | feat: render display pricing on pricing page | 1 week ago |
|
|
cc97834e5b | feat: add display pricing frontend helpers | 1 week ago |
|
|
ced4d7f323 | feat: expose model display pricing | 1 week ago |
|
|
de117ac932 | feat: add model display pricing admin api | 1 week ago |
|
|
18a3495968 | feat: add model display pricing settings | 1 week ago |
|
|
302987989b | docs: expand display pricing frontend spec | 1 week ago |
|
|
6dd8c19ea8 | docs: clarify model display pricing design | 1 week ago |
|
|
6dd9bc6dec | docs: add model display pricing design | 1 week ago |
|
|
c9988af52b |
feat: 用户创建同步到从节点及 i18n 翻译
新建用户时异步同步到海外 slave 节点,确保跨区域用户数据一致。 补充前两个 commit 全部新功能的中英文翻译字符串。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
84d364d76a |
feat: 使用日志计费详情展开与 completion_tokens 修复
前端新增 UsageLogExpandedDetail 点击展开面板和 BillingDetailSummary hover 计费详情,后端修复 per_1m_tokens 任务结算日志的 completion_tokens 从零值改为写入 total_tokens 实际值。新增 taskBillingLogMerge 合并 任务预扣/结算日志。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
2e40bab478 |
feat: 模型定价矩阵价格取中位数展示
前端模型定价配置 UI 增强(Editor/GeneratorModal/PreviewPanel/PricingGrid), 后端新增 representativeMatrixPrice 取矩阵定价表中位数作为展示价格, 替代原来的 model_price=0 导致的价格展示为空问题。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
ec7c76f221 |
merge: feat/video-pricing-table-codex into master
- feat: video multidimensional matrix pricing system - feat: doubao aiping video channel support - feat: relay capture JSON storage with header redaction - feat: SSE stream error handling (Gemini/OpenAI/Responses) - test: E2E test infrastructure (Playwright + Chromium) Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
0e3b98efe0 |
chore: commit relay SSE error handling fixes and E2E mock server
- relay: SSE stream error exposure improvements (Claude/Gemini/OpenAI) - test: add mock capture server for E2E relay capture testing Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
139b6a6075 | feat: store relay capture records as json | 2 weeks ago |
|
|
65551e8003 | refactor: prepare relay capture json helpers | 2 weeks ago |
|
|
19d55dc255 | docs: design relay capture json storage | 2 weeks ago |
|
|
1d01c4cd64 |
feat: 增强流式响应错误处理,SSE内嵌错误正确曝光为API错误
- Gemini stream: 捕获 promptFeedback.blockReason 后停止流并返回 OpenAPI 错误 - OpenAI stream: 检测 SSE 帧中的 error 响应,在首帧写入前转换为 API 错误 - Responses stream: 处理 response.error / response.failed / error 类型帧 - 新增 ClearEventStreamHeadersIfNotWritten 工具方法,错误转换时清除流式头 - 新增 streamOpenAIErrorFromData,从 SSE data 帧提取 OpenAPI 错误结构 - 上游错误体会截断后存入 UpstreamBody 供排障 Co-Authored-By: Claude <noreply@anthropic.com> |
2 weeks ago |
|
|
696c549a66 |
feat(video): add doubao aiping video pricing support
Co-Authored-By: Codex <noreply@anthropic.com> |
2 weeks ago |
|
|
99168a9660 | test: stabilize model pricing e2e runtime | 3 weeks ago |
|
|
aeff58b122 | test: stabilize model pricing generator e2e | 3 weeks ago |
|
|
5451d801f3 | test: cover model pricing e2e workflow | 3 weeks ago |
|
|
a4b824952f | test: make model pricing tab selector clickable | 3 weeks ago |
|
|
32014dde92 | test: add model pricing e2e selectors | 3 weeks ago |
|
|
eaae02e767 |
test: ignore e2e artifacts
Co-Authored-By: Codex <noreply@anthropic.com> |
3 weeks ago |
|
|
75285bcabd |
test: align playwright runtime with chromium
Co-Authored-By: Codex <noreply@anthropic.com> |
3 weeks ago |
|
|
b976b486af |
test: avoid shell by default in e2e process helper
Co-Authored-By: Codex <noreply@anthropic.com> |
3 weeks ago |
|
|
568a91dfcf |
test: harden e2e helper error handling
Co-Authored-By: Codex <noreply@anthropic.com> |
3 weeks ago |
|
|
eb285055ac | test: harden e2e process cleanup | 3 weeks ago |
|
|
32603ba385 |
docs(test): video pricing table test plan with mock server design
Co-Authored-By: Claude <noreply@anthropic.com> |
3 weeks ago |
|
|
b91d07b0da | test: add e2e process helpers | 3 weeks ago |
|
|
9837e80fac | test: add e2e scripts | 3 weeks ago |
|
|
cf2c47b826 | docs: add e2e test foundation plan | 3 weeks ago |
|
|
915efefaec | docs: add e2e test foundation design | 3 weeks ago |
|
|
3ac94e2e6f | test(task-pricing): cover model pricing api and remix snapshots | 3 weeks ago |
|
|
19bdae5509 | fix(stream): default non-positive streaming timeout | 3 weeks ago |
|
|
089bc88015 | docs(task-pricing): document video pricing table testing | 3 weeks ago |
|
|
43b0bf293d | test(task-pricing): cover video usage billing flows | 3 weeks ago |
|
|
b36edf7f43 | feat(ui): add model multidimensional pricing editor | 3 weeks ago |
|
|
aff3284d02 | feat(ui): add model pricing config utilities | 3 weeks ago |
|
|
f3778e3931 | test(task-pricing): cover doubao usage parsing | 3 weeks ago |
|
|
285323d765 | feat(task-pricing): expose model pricing rules api | 3 weeks ago |
|
|
3a52f31608 | feat(task-pricing): gate channels for matrix usage billing | 3 weeks ago |
|
|
db8eacd558 | feat(task-pricing): persist matrix billing snapshots | 3 weeks ago |
|
|
8f2b014366 | feat(task-pricing): add matrix lookup decisions | 3 weeks ago |
|
|
ae1e68792a | feat(task-pricing): resolve pricing dimensions from task requests | 3 weeks ago |
|
|
b4301a53ae | feat(task-pricing): validate and cache model pricing rules | 3 weeks ago |
|
|
95060ea0aa | feat(task-pricing): add pricing config and decision types | 3 weeks ago |
|
|
8f89828ca3 |
fix(user-migration): 禁止冲突用户非法合并
- 前端冲突处理页面禁用 cn_already_synced_to_ov 场景的 merge 操作 - 后端 resolve 接口拒绝将已存在 synced copy 的 CN 用户作为 merge 目标 - 补充对应控制器测试 Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
2f7a71d61e |
fix(user-migration): 补齐主节点保护与导入校验
- 为 user-migrations 管理路由增加仅 master 节点可访问的保护 - 为导入创建用户补充用户名、显示名、邮箱长度校验及测试 - 调整测试环境部署脚本,自动写入 CN/OV 的 NODE_TYPE Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
9aaa6baa4c |
merge: user migration selective batches into master
# Conflicts: # controller/user.go # web/src/i18n/locales/en.json # web/src/i18n/locales/zh-CN.json |
1 month ago |
|
|
65e6a82d34 |
feat(user-migration): support selective batches and cancellation
Add candidate user selection and explicit_ids batch creation flow for overseas migration. Also add cancellable batches to release selected-user locks and extend backend, frontend, and e2e coverage for the new flow. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
ca3f8765bb |
test(user-migration): complete migration coverage and e2e flow
Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
a72d76522f | fix(deploy): 部署脚本 tag 格式与 Makefile 对齐,加入 commit hash 和 dirty 标记 | 1 month ago |
|
|
de993c5605 |
feat(docker): 镜像 tag 加入 commit hash,未提交代码标记 dirty
格式:{时间}-{分支}-{短commit}[-dirty]
示例:202606051500-master-e9d5d51b(干净)
202606051500-master-e9d5d51b-dirty(有未提交改动)
|
1 month ago |
|
|
e9d5d51b37 |
fix(sync): 修复 Slave 端编辑同步用户时误报余额不可修改的问题
前端对同步用户用 SyncedQuota 替代 Quota 展示,但后端比较时 用的是数据库真实 Quota 字段,导致两者不同而触发拦截。 改为比较 SyncedQuota,只有真正修改额度时才拦截。 |
1 month ago |
|
|
2544556f5d | fix(ratio): 缓存读取倍率为 0 时前端不显示,补充 extraText 说明 | 1 month ago |
|
|
dce7631092 | docs(ratio): 缓存创建倍率 extraText 补充设置为 0 时隐藏说明 | 1 month ago |
|
|
1c15ee996d |
feat(ratio): 升级未设置倍率模型编辑弹窗,对齐可视化编辑器
- 将操作列「高级比例」按钮改为编辑图标按钮,打开完整编辑弹窗 - 编辑弹窗新增定价模式切换(按量计费/按次计费) - 按量计费下支持按倍率设置和按价格设置两种子模式 - 按倍率:模型倍率、补全倍率 - 按价格:输入价格、输出价格($/1M tokens),自动计算倍率 - 高级比例区域:缓存读取、缓存创建、图片、音频输入、音频输出 - 修复 SubmitData 遗漏高级倍率字段保存的 bug(CacheRatio 等 5 个字段) - 清理死代码(未调用的转换函数、残留 console.log) |
1 month ago |
|
|
182ab5ff83 |
refactor(ratio): move advanced ratios to edit modal in unset models page
- Remove advanced ratio columns from the table - Add 'advanced ratios' button in action column - Add edit modal with Form.InputNumber fields (same as visual editor) - Table stays clean, advanced settings only visible on click Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
1fdf0a3f47 |
feat(ratio): add advanced ratio fields to unset models editor
- Add cache/image/audio ratio columns to the table - Add advanced ratio fields to the add model modal - Update data init, submit, and addModel to handle new fields - Use formRef to read InputNumber values (same fix as visual editor) Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
d36a686f6b |
fix(ratio): preserve 0 as valid value for advanced ratio fields
- Replace || with ?? in addOrUpdateModel to keep 0 values (0 || '' = '') - Add null check in SubmitData to avoid writing undefined - Fix formRef value reading to preserve 0 vs empty distinction Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
c824381642 |
fix(ratio): read advanced ratio values from form on submit
Form.InputNumber values were not synced to currentModel state via onChange, so the advanced ratio fields were always empty when saving. Now reads values directly from formRef in onOk handler. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
7b34fbc88a |
feat(pricing): hide cache creation price when ratio is 0
- Add cacheCreationRatio !== 0 check for row data - Add hasCacheCreationPricing check for column visibility - Models with cache read but no cache creation (e.g. Gemini) no longer show empty cache creation price column Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
a8c739e466 |
feat(ui): add extraText descriptions and fix placeholder for advanced ratios
- Add extraText for cache read ratio and cache creation ratio fields - Update placeholder format to match channel pricing convention - Add missing i18n translations for extraText and new placeholders Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
192eb70862 |
feat(save): serialize and save advanced ratio fields
- Add 5 new fields to output object - Convert model data to float and populate output fields - Serialize all fields to JSON in finalOutput - Submit 7 total option fields (2 existing + 5 new) Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
91e648c530 |
feat(add): handle advanced ratios in addOrUpdateModel
- Add 5 advanced ratio fields to updatedModel object - Use empty string fallback for undefined values Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
194ebbaa56 |
feat(edit): populate advanced ratios in edit modal
- Add 5 advanced ratio fields to formRef.current.setValues - Ensure existing ratio data displays correctly in edit modal Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
2ca560ccef |
feat(ui): add advanced ratios form section
- Add Form.Section with 5 advanced ratio input fields
- Configure min={0}, step={0.01} for all fields
- Add placeholders showing default values
- Only show when pricingMode === 'per-token'
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
|
1 month ago |
|
|
0ea023037a |
feat(ratio): parse advanced ratio fields in data initialization
- Parse 5 new JSON fields from props.options - Extend modelNames Set to include keys from all ratio fields - Add 5 new fields to model data objects Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
dab78ff2c4 |
feat(i18n): add advanced ratio translations
- Add Chinese translations for 5 advanced ratio fields - Add English translations for 5 advanced ratio fields Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
a9fc624ff2 |
docs: add extended visual ratio settings implementation plan
- 9 detailed tasks covering i18n, data, UI, and testing - Each task broken into 2-5 minute steps with exact code - Complete validation and troubleshooting sections - Estimated 2.5 hours total implementation time Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
b699dda2ad |
docs: add extended visual ratio settings design spec
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
491e9b8ee4 |
fix(pricing): GetCompletionRatio 优先使用用户设置值而非硬编码默认值
修复 GetCompletionRatio 中硬编码默认值优先于用户设置值的 bug。 之前硬编码的 gpt-5 前缀匹配 (return 8, true) 会直接返回, 导致用户在后台设置的 completion_ratio 被忽略。 现在改为先查 completionRatioMap(用户设置),再查硬编码默认值。 |
1 month ago |
|
|
7af60d3214 |
feat(pricing): 分组价格表格添加缓存读取和缓存创建价格列
- calculateModelPrice 新增缓存读取/创建价格计算逻辑 - ModelPricingTable 条件显示缓存价格列(仅 cache_ratio !== 1 时) - deploy.sh 添加镜像拉取超时(120s)和3次重试机制 |
1 month ago |
|
|
3a78e9c93a | fix(user-migration): revalidate drift and verify synced quota | 1 month ago |
|
|
36b47379f7 | test(user-migration): cover rescan and batch status regressions | 1 month ago |
|
|
a406c94b93 | feat(user-migration): add root migration dashboard | 1 month ago |
|
|
7282142316 | feat(user-migration): add root admin migration API | 1 month ago |
|
|
2ecaa4f818 | feat(user-migration): add executor and verifier | 1 month ago |
|
|
9e72c2f068 | feat(user-migration): add scan service with conflict analysis | 1 month ago |
|
|
102b15ad41 | feat(user-migration): add internal migration endpoints and client | 1 month ago |
|
|
38ba0dd24b | feat(user-migration): add imported user and oauth migration helpers | 1 month ago |
|
|
26cb2273a4 | feat(user-migration): add batch item and quota grant models | 1 month ago |
|
|
87a0f94975 |
refactor: remove user-channel-ratio feature
删除用户渠道倍率(UserChannelRatio)功能,清理全链路相关代码: 后端删除 model/controller/router 层实现,前端删除 UI 组件和计费展示逻辑。 同时恢复 Token 分组选择和价格侧边栏分组过滤器的注释代码。 Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
0c279e6662 | docs: add overseas user migration design | 1 month ago |
|
|
6ad0524172 |
test: add comprehensive pricing and channel selection tests
Add unit tests for group ratio, cross-group retry, and price calculation: - Group ratio CRUD and special ratio override (12 tests) - HandleGroupRatio with auto-group and special ratios (5 tests) - ModelPriceHelper with group ratio in price/ratio modes (3 tests) - Channel selection priority, groups, weighted LB (6 tests) - RetryParam and AutoGroup service logic (7 tests) Also restore the group pricing card in model detail SideSheet. |
1 month ago |
|
|
052b41562a |
fix(pricing): restore group info on model cards
Uncomment the group ratio display that was hidden in
|
1 month ago |
|
|
0c2ea777f2 |
fix: remove channel-pricing API calls from frontend
Backend endpoints were removed but frontend still called them causing 404s. ChannelPricingCard no longer fetches, channelPricingApi and pricingTagApi return empty stubs to keep dependent components stable. |
1 month ago |
|
|
c64e9c0b40 |
chore: remove .agents from git tracking
Add .agents/ to .gitignore and untrack skill files. |
1 month ago |
|
|
a88ab5914a |
chore: remove web/dist from git tracking
Revert .gitignore exceptions that accidentally tracked web/dist build artifacts. The directory should never be committed. |
1 month ago |
|
|
7a2f584317 |
fix(claude): use gjson/sjson for reliable thinking.type replacement
Replace bytes.Replace with gjson.GetBytes+sjson.SetBytes in passthrough mode to correctly handle JSON whitespace variations and avoid silently sending empty body on read errors. Also add log search count limit to prevent slow COUNT queries on large log tables. |
1 month ago |
|
|
8150993caf |
fix(admin): improve channel form inputs
Expose channel public_name in the edit modal and load a larger channel page when building model settings views. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
3aa4cdf2ec |
feat(models): add structured endpoint editor
Replace freeform endpoint JSON editing with a structured editor, keep advanced JSON mode, and make backend pricing parsing tolerate legacy array-form endpoint configs. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
481b43cfc5 |
fix(claude): send adaptive thinking type
Use adaptive thinking for Claude requests and rewrite passthrough bodies that still contain the legacy enabled type. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
23f9ed4ccd |
fix(redemption): block synced users from redeeming
Reject redemption requests from synced users, surface a localized API message, and cover both model- and controller-level paths with tests. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
2af3a9301d |
fix(pricing): remove stale channel pricing backend
Delete leftover channel pricing controllers and routes after removing channel pricing, and align relay price helpers with global pricing plus user channel ratio only. Co-Authored-By: Codex <noreply@anthropic.com> |
1 month ago |
|
|
3a6dec8a6c |
Merge branch 'feat/remove-channel-pricing'
# Conflicts: # model/pricing_test.go |
1 month ago |
|
|
6479883ddc |
fix: handle object-type arguments in Responses API stream
gpt-5.4 returns `arguments` as a JSON object for some tools (e.g. apply_patch, tool_search), but ResponsesOutput.Arguments was typed as string, causing json.Unmarshal to fail on every such SSE event. The failure silently dropped all output_item events and the response.completed event from the forwarded stream, so: - the client never received tool call content - usage/token counts could not be extracted (billed as 0) - client reported "stream closed before response.completed" Fix: change Arguments to json.RawMessage and add GetArguments() which normalises both forms to a plain string (unescapes JSON strings, returns raw bytes for objects). Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
674b96298c | refactor(model): remove channel pricing | 1 month ago |
|
|
415174a676 | test(model): switch pricing tests to global defaults | 1 month ago |
|
|
7ecedd8e57 | fix: restore clean go test baseline | 1 month ago |
|
|
3cf84775f2 |
feat: add channel metrics collector, fix search pagination and vendor sort order
- Register Prometheus ChannelCollector in GetMetrics() sync.Once block - Add MetricsMiddleware to /v1 and /v1beta relay router groups when metrics enabled - Remove public_name field from channel creation form and validation - Restore InvitationCard in topup page - Extract fetchSearchPage helper in useModelsData to eliminate ~50 lines of duplication; fix setSearching with finally block; use Axios params for URL encoding; fix missing vendor_counts update during paginated search - Fix vendor ordering in HomePricingFilters to follow vendorsMap sort_order; fix O(n²) Array.includes to use Set Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> |
1 month ago |
|
|
6e81faa698 |
perf: lazy-read body in DoApiRequest for passthrough model mapping
DoApiRequest previously called io.ReadAll on every request body to apply passthrough model mapping, even when the feature was disabled. This caused unnecessary memory pressure for large request bodies (long conversations, multipart image uploads). Split the guard logic into needsPassthroughModelMapping() so that ReadAll is only called when all conditions are met (feature enabled + passthrough mode + model mapped). Non-passthrough requests now pass the original io.Reader directly to http.NewRequest without buffering. Rewrite tests to verify combined guard + modify behavior. |
1 month ago |
|
|
6434a523a1 | feat: passthrough model mapping | 1 month ago |
|
|
4584a94620 | Merge branch 'feat/custom-nav-link' | 1 month ago |
|
|
d32050d081 | fix: restore login widget bundle and guard custom nav link | 1 month ago |
|
|
54573991a4 |
feat: add customLink configuration UI in header nav settings
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
1 month ago |
|
|
49c6755622 |
feat: add customLink to navigation hook with filtering logic
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
1 month ago |
|
|
79f8a18cef | feat(widget): 重构登录组件为无 UI 的 headless SDK | 1 month ago |
|
|
60a1403ae2 |
fix(widget): Dockerfile 添加 Widget 构建步骤
主前端构建后执行 widget 构建,确保 login-widget.js 被 embed 打包 |
1 month ago |
|
|
aa15965467 |
fix: Dockerfile 使用华为云镜像源解决构建网络问题
移除 # syntax 指令,将 Docker Hub 基础镜像替换为华为云 SWR 镜像,解决 WSL 环境下无法直连 Docker Hub 的问题。 |
1 month ago |
|
|
1c5e6d03d9 |
fix: 禁止修改同步用户余额
同步用户余额由主节点管理,从节点不应允许管理员修改其余额。 前端禁用额度输入和添加额度按钮,后端增加校验拦截。 |
1 month ago |
|
|
fe62509b03 | feat(metrics): 重试循环和错误处理中采集指标 | 1 month ago |
|
|
e51a66c5f9 | feat(metrics): postConsumeQuota 中采集 token/配额/延迟指标 | 1 month ago |
|
|
123f0d95ae | feat(metrics): 注册 /metrics 端点 | 1 month ago |
|
|
87628f5469 | feat(metrics): 渠道状态自定义 Collector | 1 month ago |
|
|
67e78f2c44 | feat(metrics): Gin 中间件 - 活跃请求、延迟、状态码 | 1 month ago |
|
|
9d4378dbb5 | feat(metrics): 错误分类函数及测试 | 1 month ago |
|
|
688c5f8666 | feat(metrics): Prometheus 指标定义和注册 | 1 month ago |
|
|
5afa4c48fc |
feat: 用户表新增注册时间字段 + 用户列表显示邮箱和注册时间
- model/user.go: 新增 CreatedAt 字段,Insert 时自动填充 - UsersColumnDefs.jsx: 用户列表新增邮箱、注册时间列 - makefile: 新增 build-widget target |
1 month ago |
|
|
0c369eba51 |
docs: 添加登录 Widget 使用指南
API 参考、完整示例、主题配置、常见问题等 |
1 month ago |
|
|
763cf86a7a |
fix(widget): 修复跨域 CORS 问题
- CORS 中间件改用 AllowOriginFunc 回显具体 origin,替代 AllowAllOrigins:* - userRoute 添加 CORS 中间件,登录接口支持跨域请求 - Widget fetch 移除 credentials: 'include',无需 Cookie |
1 month ago |
|
|
de3b359dea |
feat(widget): 登录 Widget v1 — 密码登录,Shadow DOM 隔离
- 新增 web/widget/ 独立 Vite 构建,IIFE 格式打包 - Shadow DOM 样式隔离,支持 light/dark 主题 - 调用现有 /api/user/login 接口,零后端改动 - 构建产物 web/dist/static/login-widget.js (146KB) |
1 month ago |
|
|
3c5b821584 |
docs: 添加登录 Widget 实施计划
6 个 Task:项目骨架 → 样式 → LoginForm → 入口 → 构建集成 → E2E 验证 |
1 month ago |
|
|
08b89e6252 |
docs: 添加登录 Widget 设计文档
纯前端登录表单 Widget,通过 <script> 标签嵌入到第三方页面, 调用现有 /api/user/login 接口,零后端改动。 |
1 month ago |
|
|
84eb6d6a88 |
feat: 端点格式限制功能 + 完善测试
当开启端点格式限制后,不同协议的渠道只能通过对应的端点访问: - Anthropic 渠道只能用 /v1/messages - Gemini 渠道只能用 /v1beta/models/* - OpenAI 系渠道只能用 /v1/chat/completions 等 - VertexAI/阿里云等支持多协议的渠道不受限 核心实现:EndpointFormatGuard 中间件从请求路径推导协议家族, 不依赖 gin context 中的 relay_format(因 group middleware 执行 时 context 中尚无此值)。 修复了 /v1/models/*path 和 /v1/engines/ 路径被错误归类为 openai 家族的 bug(实际是 Gemini 格式)。 单元测试 58 个,E2E 测试 15 个,覆盖: - 路径分类边界(含 Gemini relay 路径修复验证) - 渠道类型覆盖(含非 LLM 渠道透传) - 并发安全、运行时开关切换 - 错误消息 JSON 结构和内容验证 |
1 month ago |
|
|
4c76265f32 |
test: 添加渠道选择删除 E2E 测试 + 清理临时截图 + 新增辅助脚本
- 新增 E2E 测试脚本 test/e2e_channel_removal.py,覆盖 Token 解析、 已删除路由验证、请求自动分发、channel_id 忽略、定价无默认渠道等 - 新增 anthropic_cache_test.sh 缓存测试脚本 - 新增本地开发辅助脚本 (dev.ps1, start-local.ps1 等) - 清理根目录下 9 个临时测试截图 |
1 month ago |
|
|
c64afdc6c2 |
fix: 清理渠道选择删除后的代码质量问题
- router/api-router.go: 修复删除路由后的缩进异常 - model/channel_pricing.go: 清理 Insert/Update/Delete 中遗留空行 - controller/relay_test.go: newTestContext 接收 *testing.T 参数并调用 t.Helper() |
1 month ago |
|
|
631e030411 |
test: 添加渠道选择功能删除后的验证测试
验证 parseTokenKey 不再解析 channelId(sk-abc123:42 中的 :42 保留为 key 一部分),shouldRetry/shouldRetryTaskRelay 不再检查 specific_channel_id,所有请求都能正常重试。 |
1 month ago |
|
|
b49842d543 |
refactor: 清理 Distribute 中的冗余条件判断和缩进
- 删除 tautological 的 `if channel == nil` 外层判断(channel 在该位置永远为 nil) - 修正删除指定渠道分支后遗留的多余缩进(减少一级 tab) - 更新通道亲和性检查注释为简洁版 |
1 month ago |
|
|
ea0d4e5cde | refactor: 完整删除前端渠道选择和默认通道 UI | 1 month ago |
|
|
baed9754d0 |
refactor: 完整删除渠道选择功能后端代码
删除了用户通过 sk-{key}:{channelId} 指定渠道、Playground 下拉框选择渠道、
以及管理员设置默认渠道 (is_default) 的所有后端代码。
涉及修改:
- middleware/auth.go: 简化 parseTokenKey,移除 specific_channel_id 设置
- middleware/distributor.go: 移除指定渠道分发逻辑、默认通道检查、ModelRequest.ChannelId
- constant/context_key.go: 移除 ContextKeyTokenSpecificChannelId
- controller/playground_channels.go: 整个文件删除
- controller/relay.go: 移除 shouldRetry/shouldRetryTaskRelay 中的指定渠道跳过
- controller/channel.go: 移除 GetUserChannelsForBinding
- controller/channel_pricing.go: 移除 SetDefaultChannel/ClearDefaultChannel
- model/ability.go: 移除 GetModelChannelsForGroup
- model/channel.go: 移除 GetAllChannelsForBinding
- model/channel_pricing.go: 移除默认通道缓存及所有相关函数
- model/pricing.go: 移除 DefaultChannelName 字段
- router/api-router.go: 移除 4 条路由
- i18n: 移除 channel.id_format_error 和 distributor.invalid_channel_id
|
1 month ago |
|
|
547ed7e307 | docs: 渠道选择功能完整删除设计方案 | 1 month ago |
|
|
f1a68027f1 |
Revert "feat(cache): OpenAI→Claude 转换自动注入 prompt caching"
This reverts commit
|
1 month ago |
|
|
ff14caed7c |
feat(cache): OpenAI→Claude 转换自动注入 prompt caching
- ClaudeRequest 新增 CacheControl 字段
- RequestOpenAI2ClaudeMessage 自动注入顶层 cache_control: {"type":"ephemeral"}
- System 复合内容和普通消息 content block 的 CacheControl 透传
- ParseContent 新增 []MediaContent 类型断言,修复透传路径
- 新增 8 个单元测试覆盖自动注入、透传、序列化等场景
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
1 month ago |
|
|
155aa12014 |
perf(docker): BuildKit 缓存加速构建 + fix: 模型列表改用 enabled 接口
Dockerfile 添加 --mount=type=cache 持久化 bun/go 缓存, 前端改代码时跳过后端构建(~10s),后端改代码时跳过前端(~30s)。 模型限流下拉改用 /api/channel/models_enabled 只返回已启用模型。 |
1 month ago |
|
|
f07f2c25f4 |
feat(rate-limit): 用户模型限流配置改用下拉选择模型
将添加模型限流时的文本输入改为从渠道模型列表中选择, 并自动过滤已配置的模型,提升配置体验和准确性。 |
1 month ago |
|
|
9ee5f535c8 |
Merge branch 'feat/relay-capture'
Relay 请求抓包日志 + 用户模型 RPM 限流 |
1 month ago |
|
|
714d8239d1 |
feat: Relay 请求抓包日志 + 用户模型 RPM 限流
两个独立功能: 1. Relay Capture:对开启 capture_relay 的用户记录完整请求/响应到本地文件 - 单 writer goroutine + 128 槽 channel 架构,无 OOM 风险 - 支持流式 SSE 和非流式请求 - 30s 超时保护 2. User-Model RPM Rate Limiting:按 用户+模型 维度的 RPM 限流 - Redis 令牌桶 + 内存滑动窗口双路径,Redis 故障自动降级 - 请求失败时退还令牌/配额(Redis Refund Lua 脚本 + 内存 Refund) - Redis pipeline 错误处理:失败时清理缓存 key - 前端用户编辑弹窗新增模型限流配置面板 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
1 month ago |
|
|
69ed7d85d7 |
feat(sidebar): 添加监控外部链接菜单项
在管理员菜单区域新增「监控」入口,仅 root 用户可见, 点击后在新窗口打开外部监控面板。支持在侧边栏模块管理中切换显示/隐藏。 同时在 renderWrapper 中添加外部链接的通用支持。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
bae5c1faa5 |
feat(pricing): 隐藏模型广场标签和端点类型筛选器
注释掉模型广场侧边栏和筛选弹窗中的 PricingTags 和 PricingEndpointTypes 组件, 添加 EndpointType 常量及相关组件的文档注释。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
e6b4a06557 |
feat(logs): 隐藏使用日志中的分组信息展示
使用注释方式隐藏分组列和分组倍率标签,便于后续恢复: 1. 注释 UsageLogsColumnDefs.jsx 中的分组列定义 2. 注释 useUsageLogsData.jsx 中的 COLUMN_KEYS.GROUP 和默认可见性配置 3. 将计费详情中的"分组倍率"改为"倍率" 后端计费逻辑不变,仍正常应用分组倍率。 Co-Authored-By: Claude Sonnet 4 <noreply@anthropic.com> |
2 months ago |
|
|
f9eeb7e47e |
fix(billing): 修复 Claude 计费明细未应用用户渠道折扣的问题
renderClaudeModelPrice 函数缺少 userChannelRatio 参数声明,导致运行时 ReferenceError;同时价格计算未乘以用户倍率,明细文本也未显示折扣信息。 Co-Authored-By: Claude Sonnet 4 <noreply@anthropic.com> |
2 months ago |
|
|
f7ff684d0d |
feat(pricing): 渠道定价展示用户折扣价格
登录用户查看渠道定价时,按用户分组倍率和渠道个人倍率计算实际折扣价 并在价格列展示划线原价与折扣价,新增 /pricing/user/*model 接口。 Co-Authored-By: Claude Sonnet 4 <noreply@anthropic.com> |
2 months ago |
|
|
e1b6aff7e1 |
fix(log): 对普通用户隐藏 Chat ID 和 Upstream ID
后端 GetUserLogs 返回前清空 chat_id/upstream_id(omitempty 不输出), 前端搜索框和展开行详情加 isAdminUser 守卫。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
02e0940555 |
fix(playground): 清空互斥模型默认列表
初始不预填任何模型,由管理员在运营设置页面手动配置。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
2d1a8d2f28 |
fix: 设置 chat_id + 优化日志查询参数构建
- Claude 和 OpenAI Responses 转发中设置 RelayChatID - 日志查询使用 URLSearchParams 替代手动字符串拼接 - 日志查询日期范围改为可选,不再强制默认今天 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
3e86f523db |
feat(playground): temperature/top_p 参数互斥设置
部分上游模型不允许同时设置 temperature 和 top_p,新增管理员可配置的 模型前缀列表,匹配的模型在 Playground 中自动互斥切换两个参数。 - 后端新增 PlaygroundSetting 配置 + GET /api/playground/config 公开端点 - 运营设置页面新增 Playground 互斥模型前缀编辑 textarea - Playground 自动检测模型名匹配,启用一个参数自动禁用另一个 - 默认包含 deepseek-v4-flash/pro、deepseek-reasoner、o1-/o3-/o4-、gpt-5 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
d26ad9e635 |
feat: 渠道-模型级联选择 + 隐藏首页统计卡片
- UserRatioSection 模型输入改为级联下拉,选渠道后自动加载模型列表 - models 信息嵌入 channelOptions 消除冗余 channelMap 状态 - 注释掉首页 HeroSection 统计数据卡片 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
bb708510fe |
docs: 用户倍率渠道-模型级联选择器设计文档
Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
b026f6bc04 |
feat: 用户倍率前端展示 + 渠道下拉选择 + 上游调试日志
- UserRatioSection 渠道 ID 输入改为带搜索的下拉选择器 - 使用日志详情、计费过程、展开行均显示用户倍率 - renderLogContent suffix 统一追加,消除 4 处重复 - 提取 dumpUpstreamRequest 辅助函数,debug 日志走 SysLog Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
9dbf3365fe |
feat(pricing): 添加用户-模型-渠道倍率功能
为每个用户添加独立于全局定价的倍率乘数,支持按 (userId, model, channel) 三维 精确控制用户级别定价。倍率以乘法叠加在现有 modelRatio × groupRatio 之上。 - 新增 user_channel_ratios 表,写入内存缓存(写穿透模式) - 新增 /api/user_channel_ratio/ CRUD API(AdminAuth) - 计费路径集成:compatible_handler、PostClaudeConsumeQuota、 PostWssConsumeQuota、PostAudioConsumeQuota、calculateAudioQuota - 前端:EditUserModal 内嵌 UserRatioSection 组件 - 修复 Update() 缓存 key 零值 bug(先加载再更新) - 修复 PreWssConsumeQuota 缺少 UserChannelRatio 导致归零 - 修复 ModelPriceHelperPerCall 缺少默认值 1.0 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
17a0c2c26a |
refactor(playground): 隐藏对话页面的分组选择框
Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
b033edec56 |
fix: pending sync records 队列被 quota=0 旧记录阻塞
GetPendingRecordsForSync 添加 quota > 0 过滤条件,避免因上游超时 产生的 quota=0 旧记录(4月29日网络超时遗留)占满批次,阻塞新 的正常的扣费同步到 master 节点。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
8945f8c39e |
fix: synced 用户剩余额度取 synced_quota + 记录 chat ID
- 前端用户列表:synced 用户显式取 synced_quota 字段作为剩余额度 - 各 OpenAI handler 中记录 relayChatID 用于日志追踪 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
e41f2951f7 | fix: preserve relay body and capture invalid responses body | 2 months ago |
|
|
b2cb7bcf60 | feat: record chat and upstream ids in usage logs | 2 months ago |
|
|
caf457ce9e |
feat: 错误时记录上游响应体,支持流式
- 提取 TruncateBody 公共函数,截断到 2KB 避免日志过大 - 新增 handleResponsesStreamError 统一 SSE 错误提取逻辑 - 流式/非流式 Responses API 错误路径均捕获 UpstreamBody - 添加 upstream_body 和 truncate_body 单元测试 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
575f84064e |
feat: codex API key 模式 + 前端凭证回显修复
- codex adaptor 支持 API key 和 OAuth 两种认证模式 - 提取 setupOAuthHeader 和 shouldUseChatCompletionsViaResponses - 修复编辑页 codex_credential_mode 切换/回显不同步问题 - ResponsesStreamResponse 添加 Error 字段支持独立错误事件 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
ca91922904 | docs: add log chat/upstream id design | 2 months ago |
|
|
332d5a62f2 | feat: support codex api key credentials | 2 months ago |
|
|
1c0898257c | feat: add redemption remarks | 2 months ago |
|
|
f18e51e3a9 | feat: add redemption remark support | 2 months ago |
|
|
fae24fdf37 | test: stabilize channel affinity usage cache tests | 2 months ago |
|
|
fc0257f6ca | chore: ignore local worktrees | 2 months ago |
|
|
d8272e7707 |
feat(channel): 添加渠道"对外名称"(public_name)字段
为 Channel 模型新增 public_name 字段,让管理员可以为每个渠道 设置用户可见的友好名称(如"标准通道"、"高速通道"),替代前端 硬编码的"通道一/二/三"。Playground 通道选择器展示对外名称, "渠道"统一改为"通道"。新建渠道时对外名称必填,编辑时可选。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
73d10b5799 |
feat: 错误日志记录上游 request-id 和响应体,Playground 渠道路由改为 header 传递
- 错误日志新增 upstream_request_id(从 Anthropic/OpenAI 响应 header 提取)和 upstream_body(截断 2KB) - 修复 RelayErrorHandler 内部 WithOpenAIError/NewOpenAIError 分支丢失上游字段的 bug - Playground 渠道和分组改为通过 X-Channel-Id/X-Group header 传递,而非 body 字段 - Distributor 中间件支持从 header 回退读取 channel_id 和 group Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
5b8303ae81 |
fix: 登录后跳转来源页,充值金额校验优化
- LoginForm: 登录成功后跳转回登录前的页面,而非固定 /console - RechargeCard: 输入金额低于最低值时显示警告提示 - .dockerignore: 排除 .claude、.plans、脚本等非项目文件 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
8bb2883d7e |
feat(home): 首页定价卡片添加"去体验"按钮,跳转 Playground
- PricingCardView 新增 showTryButton 属性,点击跳转 Playground 并预选模型 - 移除 PricingCardView 中未使用的 props (selectedGroup, currency 等) - 首页页脚路由重命名: terms→user-agreement, usage-policy→privacy-policy - i18n: "操练场"更名为"对话" Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
a0ac37b2a9 |
fix: 替换 println 为 SysLog,修复 logger nil context 崩溃
- logger: 增加 ctx nil 检查,避免系统级日志 panic - relay: 将 DebugEnabled 下的 println 替换为 SysLog - price: 修复渠道定价在 ChannelMeta 未初始化时的获取失败问题 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
74cc8c0d56 |
feat(playground): 添加渠道选择功能,支持指定渠道体验模型
- 新增 /api/user/model_channels 接口,返回模型可用渠道及默认渠道 - Playground 设置面板添加渠道选择下拉框 - Distribute 中间件支持从请求体读取 channel_id 指定渠道 - 支持 URL 参数 ?model=xxx 直接选择模型 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
a34d44a824 |
fix(user): 修复从节点同步用户余额显示为 0 的问题
同步用户(synced)的实际余额存储在 synced_quota 字段, 但用户列表/详情接口返回的是 quota 字段(值为 0)。 新增 ApplySyncedQuota 方法,在 API 返回前将 synced_quota 赋值给 quota,使前端能正确显示余额。 同时优化 GetUserModels:map 去重替代 O(n) 线性扫描, 禁用模型查询移至 model 层,新增 StringsSubtract 工具函数。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
967309fb56 |
feat: 添加邮箱后缀注册额度规则功能
管理员可在用户管理页面配置「邮箱后缀→初始额度」映射规则, 用户注册时根据邮箱后缀自动匹配并发放对应额度(替代默认额度)。 - 新增 email_quota_rule 数据表 + CRUD API(管理员权限) - 内存缓存匹配,启动时加载,增删改时刷新 - 注册流程 Insert/InsertWithTx/FinalizeOAuthUserCreation 同步支持 - 前端用户管理页新增 Tab 展示规则管理卡片 - 含 24 个测试(model 16 + controller 8) Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
83f767b40b |
Merge branch 'worktree-feat-captcha'
# Conflicts: # web/src/components/topup/RechargeCard.jsx |
2 months ago |
|
|
d24d68a916 |
fix(i18n): 修复令牌和充值页面货币显示不一致问题
- 令牌创建:快捷选项标签根据系统货币设置动态生成,替换硬编码美元 - 充值页面:输入框显示值与预设卡片统一换算为本地货币,内部值仍为原始单位 - 使用 getQuotaPerUnit() 替代魔数 500000 - 消除 RechargeCard 中 getCurrencyConfig() 的 N+1 重复调用 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
fe88bc7758 | Merge branch 'worktree-feat-captcha' | 2 months ago |
|
|
1912dd3b72 |
feat: 添加图片验证码功能,防止脚本批量注册
- 后端:新增 GET /api/captcha 接口,使用 base64Captcha 生成图片验证码 - 后端:发送邮箱验证码时校验图片验证码(CaptchaEnabled 开关控制) - 前端:注册表单邮箱验证码前增加图片验证码输入,支持点击刷新 - 管理后台:系统设置新增「图片验证码」开关 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
20163dbad5 |
fix(legal): 法律文档页面切换语言时强制重新加载内容
给 DocumentRenderer 添加 key={lang} 属性,确保切换语言时组件
重新挂载并从 API 获取对应语言的内容,而非使用缓存内容。
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
2 months ago |
|
|
56e8d14485 |
fix: 语言切换后 TextArea 内容丢失
添加 useEffect 在 editingLang 变化时同步 form values, 确保重新挂载的 TextArea 能从 React state 中读取正确内容。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
65bd1dabb8 |
fix: 法律文档编辑器切换语言时 TextArea 内容未刷新
给每个 TextArea 添加 key={editingLang} 强制组件在语言切换时重新挂载。
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
2 months ago |
|
|
934b4b10a7 |
feat: 法律文档双语配置(中/英)+ 精简语言支持至中英双语
- LegalSettings 字段拆分为 _zh/_en 双语对 - API 支持 ?lang= 参数返回对应语言内容 - 后台设置页每个文档配语言切换按钮(中文/English) - 前端页面根据 UI 语言自动请求对应内容 - 移除 fr/ru/ja/vi/zh-TW 语言支持 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
676902ed58 |
Merge branch 'worktree-feat-terms-usage-policy'
# Conflicts: # web/src/i18n/locales/en.json # web/src/i18n/locales/fr.json # web/src/i18n/locales/ja.json # web/src/i18n/locales/ru.json # web/src/i18n/locales/vi.json # web/src/i18n/locales/zh-CN.json # web/src/i18n/locales/zh-TW.json |
2 months ago |
|
|
2757f77011 |
feat(i18n): 补全服务条款和使用政策的多语言翻译
在 zh-CN、en、zh-TW、ja、fr、ru、vi 七个语言文件中添加 服务条款和使用政策的完整翻译 key。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
65c6446a30 | Merge branch 'worktree-feat-terms-usage-policy' | 2 months ago |
|
|
1cc75aa9a2 |
feat: 服务条款和使用政策页面,支持后台 Markdown 配置
复用现有 LegalSettings 模式,新增 TermsOfService 和 UsagePolicy 字段, 添加 /terms 和 /usage-policy 前端页面及 API 端点。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
ed195cada5 |
fix(frontend): sort_order 默认值显示为"未设置",编辑弹窗增加提示
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
d5157a779c |
fix(pricing): 统一 sort_order 排序逻辑,有 meta 但未设置的不优先
之前有 meta 记录但 sort_order=999999 的模型排在没 meta 记录的模型前面, 导致部分"未设置排序"的模型仍然挤在前面。现在统一处理:不论是否有 meta 记录,sort_order=999999 的都视为未设置,排在后面。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
54e53108ec |
fix(sort): sort_order 默认值改为 999999,简化排序逻辑
未设置排序的记录 sort_order=999999 自然排在后面,无需 CASE WHEN。 - GORM 默认值 default:0 → default:999999 - 数据库迁移:将现有 sort_order=0 的记录更新为 999999 - 回退 CASE WHEN 排序逻辑,恢复简单的 sort_order ASC, id ASC - 前端编辑弹窗默认值同步改为 999999 - 表格列中 999999 显示为空(表示未设置) Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
3b99c83e32 |
fix(sort): sort_order=0 的记录排到最后,非0值按升序排列
模型和供应商查询统一排序规则:sort_order=0 视为未设置排到最后, 非0值按升序排列。涉及 GetAllModels、SearchModels、GetAllVendors、 SearchVendors 以及 pricing 的内存排序。 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
518f1ab87f |
refactor(frontend): 移除拖拽排序,改为 sort_order 数值输入
回退 @dnd-kit 拖拽实现,改为在编辑弹窗中直接设置 sort_order 数值: - 模型编辑弹窗新增「排序」InputNumber 字段 - 供应商编辑弹窗新增「排序」InputNumber 字段 - 模型表格新增 sort_order 显示列 - 移除 @dnd-kit 依赖及所有 DnD 相关代码 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
65880fc69b |
feat(frontend): Vendor Tab 和 Model 表格支持拖拽排序,移除定价页字母排序
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
5e0f38d26c |
chore(frontend): 安装 @dnd-kit 拖拽排序库
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
e507897f21 |
feat(pricing): updatePricing 按 sort_order 排序 vendors 和 models
- vendorsList 按 vendor.SortOrder 升序排列,相同则按 ID 排序 - pricingMap 按 model.SortOrder 升序排列,相同则按模型名字典序 - 修复默认通道设置的 API 错误响应未正确展示的问题 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
749064932c |
feat: 新增 PUT /api/vendors/reorder 和 /api/models/reorder 批量排序接口
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
f8859b7c35 |
feat: Vendor 和 Model 添加 sort_order 字段,查询按 sort_order ASC 排序
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
3a29f3772f |
feat(channel): 添加模型默认通道功能,支持优先路由和卡片标识
- 新增 is_default 字段和缓存层,支持管理员为模型指定默认通道 - Distribute 中间件优先级调整:Token 指定 → 默认通道 → 亲和性 → 随机 - 模型定价卡片和详情弹窗展示默认通道 amber 标识 - 管理后台定价页面新增星标切换默认通道 - 新增 set_default / clear_default API 和 6 个单元测试 - 简化卡片价格显示(移除内联缓存价格,改为详情弹窗展示) Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
93e331b624 |
feat: 支持 LOGO_FILE_PATH 环境变量指定本地 Logo,优化定价与语言设置
- 新增 LOGO_FILE_PATH 环境变量,优先级高于数据库配置,支持本地文件服务 - 渠道定价高级字段(缓存/图片/音频)不再回退全局默认值,未设置直接返回 0 - 缓存价格单位从表头移到具体价格值,新增渠道 ID 复制功能 - 修复默认语言在用户已有偏好时仍被覆盖的问题 - img 标签统一添加 referrerPolicy/crossOrigin 防止跨域问题 - 新增缓存倍率和最小余额阈值的详细说明文本 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
94f24b29f7 |
style: 首页文案"全球"改为"顶级"
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
f0f0215b20 |
fix(pricing): 修复缓存倍率写入默认值1和缓存创建token丢失
1. CopyGlobalPricing 检查 ratio 查找的 bool 返回值,未找到时写 0 而非写入 fallback 值 1,避免覆盖全局正确的 0.1 2. Claude 响应使用 GetCacheCreationTotalTokens() 替代直接读 CacheCreationInputTokens,兼容新版本子对象格式 3. 新增 GetAudioRatioV2/GetAudioCompletionRatioV2 带 bool 返回 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com> |
2 months ago |
|
|
387f6c1ae0 |
feat(pricing): 定价数据源切换到渠道表,新增缓存价格展示
- 后端 pricing API 从 channel_pricings 表获取实际定价,选取最便宜渠道 - 提取 applyGlobalDefault 辅助函数消除全局回退逻辑重复 - price.go 重构扩展比率为局部变量,简化回退逻辑 - 定价卡片新增缓存读取/创建价格,改为两行布局防止溢出 - ChannelPricingCard 缓存列显示实际价格而非倍率 - 修复移动端 hero 区域 padding 过大 - 默认标签页标题改为 Loading... - 新增缓存读取/创建 i18n 翻译(7 语言) Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
00058cd671 |
refactor(channel-pricing): 提取辅助方法消除重复代码,修复缓存一致性
- 提取 PriceData.ApplyChannelPricingRatios 消除 ModelPriceHelper 和 UpdatePriceDataForChannelPricing 中重复的 ~30 行比率回退逻辑 - 提取 ChannelPricing.ApplyFields 消除 controller 中 4 处相同的字段赋值 - 提取 setCache/removeCache 辅助函数统一写穿透缓存操作 - 将 claudeCacheCreation1hMultiplier 常量移至 types 包避免循环依赖 - 修复 BatchUpsertChannelPricing 成功后未刷新内存缓存的 bug Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
42082311d3 |
fix: 修复默认语言不生效的问题
移除 PageLayout useEffect 中多余的 localStorage 语言恢复逻辑。 i18next-browser-languagedetector 在初始化时已自动处理 localStorage, 同步的 changeLanguage 调用会覆盖异步 loadStatus 设置的管理员默认语言。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
9eb58c684e |
feat: 后台设置默认语言
管理员可在系统设置页面配置全局默认语言,未登录用户和未设置 语言偏好的已登录用户将强制使用该语言,忽略浏览器语言检测。 - common/constants.go: 新增 DefaultLanguage 变量 - model/option.go: 新增 DefaultLanguage option handler - controller/misc.go: /api/status 暴露 default_language - SystemSetting.jsx: 添加语言下拉框到通用设置 - PageLayout.jsx: 未登录用户应用默认语言 - UserContext.jsx: 无偏好用户回退到默认语言 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
a2b27ea2a6 |
docs: 后台设置默认语言实施计划
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
436fcb405b |
docs: 后台设置默认语言功能设计文档
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
8d22c4e80e |
style(home): 临时隐藏首页工具链/核心价值/工作流/生态伙伴 section
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
355757404b | merge: feat/channel-pricing-extended → master | 2 months ago |
|
|
235f7c6e5f |
refactor(channel-pricing): 提取 ParseTagIds 辅助函数 + 移除调试 console.log
- 将 tag ID 解析逻辑提取为 model.ParseTagIds,消除 model 层和 controller 层重复代码
- 统一 TrimSpace 处理(之前 controller 版本漏了)
- 移除 ChannelPricingView 残留的 console.log('Expanded keys:')
- 删除多余空行
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
|
2 months ago |
|
|
f3577590bc |
test(channel-pricing): 添加 API 端到端集成测试脚本
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
129e17ac42 |
feat(channel-pricing): 前端支持缓存/图片/音频倍率编辑和展示
- ChannelPricingView.jsx: initValues 添加 5 个高级比例字段 + 按量计费模式下新增高级比例表单区块 - ChannelPricingCard.jsx: tableData 映射 cache_ratio/cache_creation_ratio + 条件列展示 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
c84c26b5a2 |
feat(channel-pricing): GetChannelPricingByModelWithChannelInfo 支持 CASE WHEN 回退 + 扩展字段
在 SQL 查询中使用 CASE WHEN > 0 回退策略,让渠道定价的扩展比率 (cache_ratio, cache_creation_ratio, image_ratio, audio_ratio, audio_completion_ratio) 在未设置时自动回退到全局默认值, 而不是像 COALESCE 那样被零值拦截。同时更新 GROUP BY 子句包含所有 新选择的列,确保 MySQL 兼容。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
d5d4714908 |
test(channel-pricing): 缓存写穿 + 字段默认值单元测试
验证 ChannelPricing 的 Insert/Update/Delete 操作正确更新内存缓存, 以及未设置的扩展字段(CacheRatio 等)默认为零值。 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
646f37dd4b |
feat(channel-pricing): Controller 扩展 API 支持新字段 + 输入校验 + 操作日志
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
3a24fd598f |
feat(channel-pricing): ModelPriceHelper + UpdatePriceDataForChannelPricing 适配新签名,支持扩展比率覆盖
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
e5f029f91e |
feat(channel-pricing): InitDB 启动时全量加载渠道定价缓存
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
2d4c73d3aa |
refactor(channel-pricing): 结构体新增扩展字段 + 缓存重写为全量加载写穿
- ChannelPricing 结构体新增 5 个扩展计费字段: cache_ratio, cache_creation_ratio, image_ratio, audio_ratio, audio_completion_ratio - 缓存机制从 TTL+惰性加载重写为全量加载+写穿模式,消除 DB 查询延迟 - Insert/Update/Delete 改为写穿缓存(直接更新内存,不再整表失效) - 新增 LoadChannelPricingCache 启动时全量加载函数 - 删除 RefreshChannelPricingCache/InvalidateChannelPricingCache(不再需要) - BatchUpsertChannelPricing DoUpdates 列表同步新增 5 字段 - GetEffectivePricing 签名改为 (*ChannelPricing, bool)(调用方适配在后续 Task) Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> |
2 months ago |
|
|
156618fdad |
merge: feat/alipay-payment → master
支付宝当面付扫码支付集成 + 代码重构 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
|
|
b514a0798b |
refactor(payment): 合并微信/支付宝支付重复代码,删除冗余文件
- 提取通用 calcPayMoney/calcMinTopup 函数,消除微信和支付宝控制器中的重复计费逻辑 - 合并 RechargeAlipay/RechargeWechat 为 rechargeByQRCodePayment 内部函数,删除 model/topup_alipay.go - 复用已有的 wrapAsPEM 替换支付宝专用的 wrapAlipayPublicKey - 删除被 QRCodePayModal 替代的 WechatPayQRCodeModal.jsx - 修复 controller/topup.go 中支付宝代码块的缩进错误 净减 264 行代码。 Co-Authored-By: Claude <noreply@anthropic.com> |
2 months ago |
| @@ -7,4 +7,24 @@ Makefile | |||
| docs | |||
| .eslintcache | |||
| .gocache | |||
| /web/node_modules | |||
| /web/node_modules | |||
| .claude | |||
| .plans | |||
| current-page* | |||
| models-page* | |||
| pricing-page.png | |||
| login-page | |||
| scripts | |||
| relay/helper/price_test.go | |||
| .superpowers | |||
| *.png | |||
| *.bak | |||
| .worktrees | |||
| **/node_modules | |||
| **/.gocache | |||
| **/.gocache-temp | |||
| logs | |||
| *.db | |||
| *.db-journal | |||
| *.zip | |||
| web/dist | |||
| @@ -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 | |||
| @@ -36,3 +36,4 @@ | |||
| # ============================================ | |||
| # Mark web frontend as vendored so GitHub recognizes this as a Go project | |||
| electron/** linguist-vendored | |||
| .dockerignore text eol=lf | |||
| @@ -23,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/ | |||
| @@ -0,0 +1,7 @@ | |||
| # backend-dev - 发现索引 | |||
| > 纯索引——每个条目应简短(Status + Report 链接 + Summary)。 | |||
| --- | |||
| <初始为空,工作中填写> | |||
| @@ -0,0 +1,7 @@ | |||
| # backend-dev - 工作日志 | |||
| > 用于上下文恢复。压缩/重启后先读此文件。 | |||
| --- | |||
| <初始为空,工作中填写> | |||
| @@ -0,0 +1,7 @@ | |||
| # 支付宝后端 - 发现记录 | |||
| > 此任务开发中的技术发现。 | |||
| --- | |||
| <初始为空> | |||
| @@ -0,0 +1,7 @@ | |||
| # 支付宝后端 - 工作日志 | |||
| > 上下文恢复时只需读此文件。 | |||
| --- | |||
| <初始为空> | |||
| @@ -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 已在项目中 | |||
| @@ -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 审查 | |||
| @@ -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 角色(全栈) | |||
| @@ -0,0 +1,7 @@ | |||
| # alipay-payment - 发现与技术记录 | |||
| > 由团队智能体自动更新。每条标注来源。 | |||
| --- | |||
| <工作中添加条目> | |||
| @@ -0,0 +1,7 @@ | |||
| # frontend-dev - 发现索引 | |||
| > 纯索引——每个条目应简短。 | |||
| --- | |||
| <初始为空,工作中填写> | |||
| @@ -0,0 +1,7 @@ | |||
| # frontend-dev - 工作日志 | |||
| > 用于上下文恢复。 | |||
| --- | |||
| <初始为空,工作中填写> | |||
| @@ -0,0 +1,7 @@ | |||
| # 支付宝前端 - 发现记录 | |||
| > 此任务开发中的技术发现。 | |||
| --- | |||
| <初始为空> | |||
| @@ -0,0 +1,7 @@ | |||
| # 支付宝前端 - 工作日志 | |||
| > 上下文恢复时只需读此文件。 | |||
| --- | |||
| <初始为空> | |||
| @@ -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 遵循现有命名规范 | |||
| @@ -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 审查 | |||
| @@ -0,0 +1,16 @@ | |||
| # alipay-payment - 进度日志 | |||
| > 按时间线记录。每条记录谁做了什么。 | |||
| --- | |||
| ## 2026-04-14 Session 1 — 团队搭建 | |||
| ### 已完成 | |||
| - [x] 读取设计文档 | |||
| - [x] 确认团队配置:3 角色(backend-dev, frontend-dev, reviewer) | |||
| - [x] 创建规划文件和目录结构 | |||
| ### 待办 | |||
| - [ ] 启动团队成员 | |||
| - [ ] 开始并行开发 | |||
| @@ -0,0 +1,7 @@ | |||
| # reviewer - 发现索引 | |||
| > 纯索引。 | |||
| --- | |||
| <初始为空,工作中填写> | |||
| @@ -0,0 +1,5 @@ | |||
| # reviewer - 工作日志 | |||
| --- | |||
| <初始为空> | |||
| @@ -0,0 +1,16 @@ | |||
| # reviewer - 任务计划 | |||
| > 角色: 代码审查 | |||
| > 状态: pending | |||
| > 分配的任务: 等待 backend-dev 和 frontend-dev 完成后进行代码审查 | |||
| ## 任务 | |||
| - [ ] 审查 backend-dev 支付宝后端代码(安全 + 质量) | |||
| - [ ] 审查 frontend-dev 支付宝前端代码(质量 + 体验) | |||
| ## 备注 | |||
| - 支付涉及资金安全,重点关注:签名验签、金额处理、幂等性、并发安全 | |||
| - 后端审查重点:回调验签、金额单位(元 vs 分)、订单幂等 | |||
| - 前端审查重点:二维码模态框重构兼容性、支付状态轮询 | |||
| @@ -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(并行开发)。 | |||
| @@ -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 的渠道 → 显示"通道一/二/三..." | |||
| @@ -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> | |||
| @@ -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> | |||
| @@ -0,0 +1 @@ | |||
| {"reason":"idle timeout","timestamp":1777286270028} | |||
| @@ -0,0 +1 @@ | |||
| 6600 | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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" | |||
| } | |||
| } | |||
| @@ -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: | |||
| @@ -0,0 +1,30 @@ | |||
| package common | |||
| import ( | |||
| "github.com/mojocn/base64Captcha" | |||
| ) | |||
| var captchaStore = base64Captcha.DefaultMemStore | |||
| var captchaDriver = base64Captcha.NewDriverString( | |||
| 40, // height | |||
| 120, // width | |||
| 0, // noise count (auto) | |||
| base64Captcha.OptionShowSlimeLine, // show slime lines | |||
| 5, // code length | |||
| base64Captcha.TxtSimpleCharaters, // digits+letters excluding confusing chars | |||
| nil, // bg color (auto) | |||
| nil, // font storage (auto) | |||
| []string{"wqy-microhei.ttc"}, // font files | |||
| ) | |||
| var captchaInstance = base64Captcha.NewCaptcha(captchaDriver, captchaStore) | |||
| func GenerateCaptcha() (string, string, error) { | |||
| id, b64s, _, err := captchaInstance.Generate() | |||
| return id, b64s, err | |||
| } | |||
| func VerifyCaptcha(id, code string) bool { | |||
| return captchaStore.Verify(id, code, true) | |||
| } | |||
| @@ -0,0 +1,58 @@ | |||
| package common | |||
| import ( | |||
| "errors" | |||
| "strings" | |||
| ) | |||
| type CodexCredentialMode string | |||
| const ( | |||
| CodexCredentialModeAPIKey CodexCredentialMode = "api_key" | |||
| CodexCredentialModeOAuth CodexCredentialMode = "oauth" | |||
| ) | |||
| var ( | |||
| ErrCodexOAuthCredentialRequired = errors.New("codex channel: oauth credential required") | |||
| ErrCodexOAuthCredentialInvalidJSON = errors.New("codex channel: invalid oauth key json") | |||
| ErrCodexOAuthAccessTokenRequired = errors.New("codex channel: access_token is required") | |||
| ErrCodexOAuthAccountIDRequired = errors.New("codex channel: account_id is required") | |||
| ) | |||
| type CodexOAuthCredential struct { | |||
| IDToken string `json:"id_token,omitempty"` | |||
| AccessToken string `json:"access_token,omitempty"` | |||
| RefreshToken string `json:"refresh_token,omitempty"` | |||
| AccountID string `json:"account_id,omitempty"` | |||
| LastRefresh string `json:"last_refresh,omitempty"` | |||
| Email string `json:"email,omitempty"` | |||
| Type string `json:"type,omitempty"` | |||
| Expired string `json:"expired,omitempty"` | |||
| } | |||
| func ParseCodexOAuthCredential(raw string) (*CodexOAuthCredential, error) { | |||
| trimmed := strings.TrimSpace(raw) | |||
| if trimmed == "" || !strings.HasPrefix(trimmed, "{") { | |||
| return nil, ErrCodexOAuthCredentialRequired | |||
| } | |||
| var credential CodexOAuthCredential | |||
| if err := Unmarshal([]byte(trimmed), &credential); err != nil { | |||
| return nil, ErrCodexOAuthCredentialInvalidJSON | |||
| } | |||
| if strings.TrimSpace(credential.AccessToken) == "" { | |||
| return nil, ErrCodexOAuthAccessTokenRequired | |||
| } | |||
| if strings.TrimSpace(credential.AccountID) == "" { | |||
| return nil, ErrCodexOAuthAccountIDRequired | |||
| } | |||
| return &credential, nil | |||
| } | |||
| func DetectCodexCredentialMode(raw string) CodexCredentialMode { | |||
| if _, err := ParseCodexOAuthCredential(raw); err == nil { | |||
| return CodexCredentialModeOAuth | |||
| } | |||
| return CodexCredentialModeAPIKey | |||
| } | |||
| @@ -15,7 +15,16 @@ var Version = "v0.0.0" // this hard coding will be replaced automatic | |||
| var SystemName = "New API" | |||
| var Footer = "" | |||
| var Logo = "" | |||
| var LogoFilePath = "" // LOGO_FILE_PATH 环境变量指定的本地 Logo 文件路径 | |||
| func GetEffectiveLogo() string { | |||
| if LogoFilePath != "" { | |||
| return "/logo.png" | |||
| } | |||
| return Logo | |||
| } | |||
| var TopUpLink = "" | |||
| var DefaultLanguage = "" // admin-configured default language; empty = follow browser detection | |||
| // var ChatLink = "" | |||
| // var ChatLink2 = "" | |||
| @@ -49,6 +58,7 @@ var LinuxDOOAuthEnabled = false | |||
| var WeChatAuthEnabled = false | |||
| var TelegramOAuthEnabled = false | |||
| var TurnstileCheckEnabled = false | |||
| var CaptchaEnabled = false | |||
| var RegisterEnabled = true | |||
| var EmailDomainRestrictionEnabled = false // 是否启用邮箱域名限制 | |||
| @@ -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 返回指定端点类型的默认信息以及是否存在 | |||
| @@ -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} | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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 | |||
| } | |||
| } | |||
| } | |||
| @@ -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") | |||
| } | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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{}) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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() | |||
| } | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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)) | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| } | |||
| @@ -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 { | |||
| @@ -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" | |||
| @@ -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" | |||
| @@ -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) { | |||
| @@ -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, | |||
| }) | |||
| } | |||
| @@ -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, | |||
| }) | |||
| } | |||
| @@ -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), | |||
| ) | |||
| } | |||
| @@ -0,0 +1,111 @@ | |||
| package controller | |||
| import ( | |||
| "bytes" | |||
| "net/http" | |||
| "net/http/httptest" | |||
| "testing" | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/constant" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/gin-gonic/gin" | |||
| "github.com/glebarez/sqlite" | |||
| "github.com/stretchr/testify/assert" | |||
| "github.com/stretchr/testify/require" | |||
| "gorm.io/gorm" | |||
| ) | |||
| func setupCodexChannelDB(t *testing.T, key string) *gorm.DB { | |||
| t.Helper() | |||
| db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) | |||
| require.NoError(t, err) | |||
| sqlDB, _ := db.DB() | |||
| sqlDB.SetMaxOpenConns(1) | |||
| origDB := model.DB | |||
| origLogDB := model.LOG_DB | |||
| model.DB = db | |||
| model.LOG_DB = db | |||
| common.UsingSQLite = true | |||
| common.RedisEnabled = false | |||
| require.NoError(t, db.AutoMigrate(&model.Channel{})) | |||
| require.NoError(t, db.Create(&model.Channel{ | |||
| Id: 1, | |||
| Name: "codex-channel", | |||
| PublicName: "codex-channel", | |||
| Type: constant.ChannelTypeCodex, | |||
| Key: key, | |||
| Status: common.ChannelStatusEnabled, | |||
| }).Error) | |||
| t.Cleanup(func() { | |||
| model.DB = origDB | |||
| model.LOG_DB = origLogDB | |||
| _ = sqlDB.Close() | |||
| }) | |||
| return db | |||
| } | |||
| func setupCodexChannelRouter(t *testing.T, key string) *gin.Engine { | |||
| t.Helper() | |||
| setupCodexChannelDB(t, key) | |||
| gin.SetMode(gin.TestMode) | |||
| r := gin.New() | |||
| g := r.Group("/api/channel") | |||
| g.POST("/:id/codex/refresh", RefreshCodexChannelCredential) | |||
| g.GET("/:id/codex/usage", GetCodexChannelUsage) | |||
| return r | |||
| } | |||
| func TestValidateChannelAcceptsCodexAPIKey(t *testing.T) { | |||
| channel := &model.Channel{ | |||
| Name: "codex-api-key", | |||
| PublicName: "codex-api-key", | |||
| Type: constant.ChannelTypeCodex, | |||
| Key: "sk-codex-api-key", | |||
| } | |||
| require.NoError(t, validateChannel(channel, true)) | |||
| } | |||
| func TestValidateChannelRejectsCodexOAuthWithoutAccountID(t *testing.T) { | |||
| channel := &model.Channel{ | |||
| Name: "codex-oauth", | |||
| PublicName: "codex-oauth", | |||
| Type: constant.ChannelTypeCodex, | |||
| Key: `{"access_token":"token-only"}`, | |||
| } | |||
| err := validateChannel(channel, true) | |||
| require.Error(t, err) | |||
| assert.Contains(t, err.Error(), "account_id") | |||
| } | |||
| func TestRefreshCodexChannelCredentialRejectsAPIKeyMode(t *testing.T) { | |||
| router := setupCodexChannelRouter(t, "sk-codex-api-key") | |||
| req := httptest.NewRequest(http.MethodPost, "/api/channel/1/codex/refresh", bytes.NewReader([]byte(`{}`))) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| w := httptest.NewRecorder() | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| assert.Contains(t, w.Body.String(), "当前凭证方式不支持刷新凭证") | |||
| } | |||
| func TestGetCodexChannelUsageRejectsAPIKeyMode(t *testing.T) { | |||
| router := setupCodexChannelRouter(t, "sk-codex-api-key") | |||
| req := httptest.NewRequest(http.MethodGet, "/api/channel/1/codex/usage", nil) | |||
| w := httptest.NewRecorder() | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| assert.Contains(t, w.Body.String(), "当前凭证方式不支持查看用量") | |||
| } | |||
| @@ -2,6 +2,7 @@ package controller | |||
| import ( | |||
| "context" | |||
| "errors" | |||
| "fmt" | |||
| "net/http" | |||
| "strconv" | |||
| @@ -11,7 +12,6 @@ import ( | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/constant" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/QuantumNous/new-api/relay/channel/codex" | |||
| "github.com/QuantumNous/new-api/service" | |||
| "github.com/gin-gonic/gin" | |||
| @@ -42,22 +42,19 @@ func GetCodexChannelUsage(c *gin.Context) { | |||
| return | |||
| } | |||
| oauthKey, err := codex.ParseOAuthKey(strings.TrimSpace(ch.Key)) | |||
| oauthKey, err := common.ParseCodexOAuthCredential(strings.TrimSpace(ch.Key)) | |||
| if err != nil { | |||
| if errors.Is(err, common.ErrCodexOAuthCredentialRequired) { | |||
| c.JSON(http.StatusOK, gin.H{"success": false, "message": "当前凭证方式不支持查看用量"}) | |||
| return | |||
| } | |||
| common.SysError("failed to parse oauth key: " + err.Error()) | |||
| c.JSON(http.StatusOK, gin.H{"success": false, "message": "解析凭证失败,请检查渠道配置"}) | |||
| return | |||
| } | |||
| accessToken := strings.TrimSpace(oauthKey.AccessToken) | |||
| accountID := strings.TrimSpace(oauthKey.AccountID) | |||
| if accessToken == "" { | |||
| c.JSON(http.StatusOK, gin.H{"success": false, "message": "codex channel: access_token is required"}) | |||
| return | |||
| } | |||
| if accountID == "" { | |||
| c.JSON(http.StatusOK, gin.H{"success": false, "message": "codex channel: account_id is required"}) | |||
| return | |||
| } | |||
| client, err := service.NewProxyHttpClient(ch.GetSetting().Proxy) | |||
| if err != nil { | |||
| @@ -98,6 +95,7 @@ func GetCodexChannelUsage(c *gin.Context) { | |||
| ctx2, cancel2 := context.WithTimeout(c.Request.Context(), 15*time.Second) | |||
| defer cancel2() | |||
| statusCode, body, err = service.FetchCodexWhamUsage(ctx2, client, ch.GetBaseURL(), oauthKey.AccessToken, accountID) | |||
| if err != nil { | |||
| common.SysError("failed to fetch codex usage after refresh: " + err.Error()) | |||
| @@ -0,0 +1,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) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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())) | |||
| } | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -0,0 +1,128 @@ | |||
| package controller | |||
| import ( | |||
| "strconv" | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/gin-gonic/gin" | |||
| ) | |||
| func GetAllEmailQuotaRules(c *gin.Context) { | |||
| list, err := model.GetAllEmailQuotaRules() | |||
| if err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| common.ApiSuccess(c, list) | |||
| } | |||
| type CreateEmailQuotaRuleRequest struct { | |||
| EmailSuffix string `json:"email_suffix" binding:"required"` | |||
| Quota int64 `json:"quota" binding:"required"` | |||
| Enabled *bool `json:"enabled"` | |||
| Description string `json:"description"` | |||
| } | |||
| func CreateEmailQuotaRule(c *gin.Context) { | |||
| var req CreateEmailQuotaRuleRequest | |||
| if err := c.ShouldBindJSON(&req); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| existing, _ := model.GetEmailQuotaRuleBySuffix(req.EmailSuffix) | |||
| if existing != nil { | |||
| common.ApiErrorMsg(c, "该邮箱后缀已存在") | |||
| return | |||
| } | |||
| enabled := true | |||
| if req.Enabled != nil { | |||
| enabled = *req.Enabled | |||
| } | |||
| rule := &model.EmailQuotaRule{ | |||
| EmailSuffix: req.EmailSuffix, | |||
| Quota: req.Quota, | |||
| Enabled: enabled, | |||
| Description: req.Description, | |||
| } | |||
| if err := rule.Insert(); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| common.ApiSuccess(c, rule) | |||
| } | |||
| type UpdateEmailQuotaRuleRequest struct { | |||
| EmailSuffix string `json:"email_suffix"` | |||
| Quota *int64 `json:"quota"` | |||
| Enabled *bool `json:"enabled"` | |||
| Description *string `json:"description"` | |||
| } | |||
| func UpdateEmailQuotaRule(c *gin.Context) { | |||
| idStr := c.Param("id") | |||
| id, err := strconv.Atoi(idStr) | |||
| if err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| var req UpdateEmailQuotaRuleRequest | |||
| if err := c.ShouldBindJSON(&req); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| rule, err := model.GetEmailQuotaRuleById(id) | |||
| if err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| if req.EmailSuffix != "" && req.EmailSuffix != rule.EmailSuffix { | |||
| existing, _ := model.GetEmailQuotaRuleBySuffix(req.EmailSuffix) | |||
| if existing != nil && existing.Id != id { | |||
| common.ApiErrorMsg(c, "该邮箱后缀已存在") | |||
| return | |||
| } | |||
| rule.EmailSuffix = req.EmailSuffix | |||
| } | |||
| if req.Quota != nil { | |||
| rule.Quota = *req.Quota | |||
| } | |||
| if req.Enabled != nil { | |||
| rule.Enabled = *req.Enabled | |||
| } | |||
| if req.Description != nil { | |||
| rule.Description = *req.Description | |||
| } | |||
| if err := rule.Update(); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| common.ApiSuccess(c, rule) | |||
| } | |||
| func DeleteEmailQuotaRule(c *gin.Context) { | |||
| idStr := c.Param("id") | |||
| id, err := strconv.Atoi(idStr) | |||
| if err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| rule := &model.EmailQuotaRule{Id: id} | |||
| if err := rule.Delete(); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| common.ApiSuccess(c, nil) | |||
| } | |||
| @@ -0,0 +1,264 @@ | |||
| package controller | |||
| import ( | |||
| "bytes" | |||
| "encoding/json" | |||
| "net/http" | |||
| "net/http/httptest" | |||
| "strconv" | |||
| "testing" | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/glebarez/sqlite" | |||
| "github.com/gin-gonic/gin" | |||
| "github.com/stretchr/testify/assert" | |||
| "github.com/stretchr/testify/require" | |||
| "gorm.io/gorm" | |||
| ) | |||
| func setupEmailQuotaRuleControllerDB(t *testing.T) *gorm.DB { | |||
| t.Helper() | |||
| db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) | |||
| require.NoError(t, err) | |||
| sqlDB, _ := db.DB() | |||
| sqlDB.SetMaxOpenConns(1) | |||
| origDB := model.DB | |||
| model.DB = db | |||
| common.UsingSQLite = true | |||
| common.RedisEnabled = false | |||
| require.NoError(t, db.AutoMigrate(&model.EmailQuotaRule{})) | |||
| t.Cleanup(func() { | |||
| model.DB = origDB | |||
| sqlDB.Close() | |||
| }) | |||
| return db | |||
| } | |||
| func setupEmailQuotaRuleRouter() *gin.Engine { | |||
| gin.SetMode(gin.TestMode) | |||
| r := gin.New() | |||
| g := r.Group("/api/email_quota_rule") | |||
| { | |||
| g.GET("/", GetAllEmailQuotaRules) | |||
| g.POST("/", CreateEmailQuotaRule) | |||
| g.PUT("/:id", UpdateEmailQuotaRule) | |||
| g.DELETE("/:id", DeleteEmailQuotaRule) | |||
| } | |||
| return r | |||
| } | |||
| func TestGetAllEmailQuotaRules_Empty(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("GET", "/api/email_quota_rule/", nil) | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.True(t, resp["success"].(bool)) | |||
| assert.Empty(t, resp["data"]) | |||
| } | |||
| func TestCreateEmailQuotaRule_Success(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "email_suffix": "@test.com", | |||
| "quota": 500000, | |||
| "description": "Test company", | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.True(t, resp["success"].(bool)) | |||
| data := resp["data"].(map[string]interface{}) | |||
| assert.Equal(t, "@test.com", data["email_suffix"]) | |||
| assert.Equal(t, float64(500000), data["quota"]) | |||
| assert.Equal(t, true, data["enabled"]) | |||
| } | |||
| func TestCreateEmailQuotaRule_Duplicate(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "email_suffix": "@dup.com", | |||
| "quota": 100, | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var firstResp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &firstResp)) | |||
| assert.True(t, firstResp["success"].(bool)) | |||
| // Second create should fail | |||
| w2 := httptest.NewRecorder() | |||
| req2, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body)) | |||
| req2.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w2, req2) | |||
| var resp2 map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w2.Body.Bytes(), &resp2)) | |||
| assert.False(t, resp2["success"].(bool)) | |||
| } | |||
| func TestCreateEmailQuotaRule_MissingFields(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "description": "no suffix", | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.False(t, resp["success"].(bool)) | |||
| } | |||
| func TestUpdateEmailQuotaRule_Success(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| // Create first | |||
| rule := &model.EmailQuotaRule{EmailSuffix: "@up.com", Quota: 100, Enabled: true} | |||
| require.NoError(t, rule.Insert()) | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "quota": 999, | |||
| "description": "updated desc", | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("PUT", "/api/email_quota_rule/"+itoa(rule.Id), bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.True(t, resp["success"].(bool)) | |||
| data := resp["data"].(map[string]interface{}) | |||
| assert.Equal(t, float64(999), data["quota"]) | |||
| assert.Equal(t, "updated desc", data["description"]) | |||
| } | |||
| func TestUpdateEmailQuotaRule_ToggleEnabled(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| rule := &model.EmailQuotaRule{EmailSuffix: "@toggle.com", Quota: 500, Enabled: true} | |||
| require.NoError(t, rule.Insert()) | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "enabled": false, | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("PUT", "/api/email_quota_rule/"+itoa(rule.Id), bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.True(t, resp["success"].(bool)) | |||
| data := resp["data"].(map[string]interface{}) | |||
| assert.Equal(t, false, data["enabled"]) | |||
| // Cache should reflect disabled | |||
| assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@toggle.com")) | |||
| } | |||
| func TestDeleteEmailQuotaRule_Success(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| rule := &model.EmailQuotaRule{EmailSuffix: "@del.com", Quota: 100, Enabled: true} | |||
| require.NoError(t, rule.Insert()) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("DELETE", "/api/email_quota_rule/"+itoa(rule.Id), nil) | |||
| router.ServeHTTP(w, req) | |||
| var resp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) | |||
| assert.True(t, resp["success"].(bool)) | |||
| // Verify deleted | |||
| assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@del.com")) | |||
| } | |||
| func TestCRUD_FullFlow(t *testing.T) { | |||
| setupEmailQuotaRuleControllerDB(t) | |||
| router := setupEmailQuotaRuleRouter() | |||
| // 1. Create | |||
| body, _ := json.Marshal(map[string]interface{}{ | |||
| "email_suffix": "@full.com", | |||
| "quota": 1000, | |||
| "description": "full flow test", | |||
| }) | |||
| w := httptest.NewRecorder() | |||
| req, _ := http.NewRequest("POST", "/api/email_quota_rule/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var createResp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &createResp)) | |||
| assert.True(t, createResp["success"].(bool)) | |||
| // 2. List | |||
| w = httptest.NewRecorder() | |||
| req, _ = http.NewRequest("GET", "/api/email_quota_rule/", nil) | |||
| router.ServeHTTP(w, req) | |||
| var listResp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &listResp)) | |||
| data := listResp["data"].([]interface{}) | |||
| assert.Len(t, data, 1) | |||
| // 3. Update | |||
| ruleId := itoa(int(data[0].(map[string]interface{})["id"].(float64))) | |||
| body, _ = json.Marshal(map[string]interface{}{"quota": 2000}) | |||
| w = httptest.NewRecorder() | |||
| req, _ = http.NewRequest("PUT", "/api/email_quota_rule/"+ruleId, bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| router.ServeHTTP(w, req) | |||
| var updateResp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &updateResp)) | |||
| assert.True(t, updateResp["success"].(bool)) | |||
| // 4. Verify cache | |||
| assert.Equal(t, int64(2000), model.MatchEmailQuotaRule("user@full.com")) | |||
| // 5. Delete | |||
| w = httptest.NewRecorder() | |||
| req, _ = http.NewRequest("DELETE", "/api/email_quota_rule/"+ruleId, nil) | |||
| router.ServeHTTP(w, req) | |||
| var delResp map[string]interface{} | |||
| require.NoError(t, json.Unmarshal(w.Body.Bytes(), &delResp)) | |||
| assert.True(t, delResp["success"].(bool)) | |||
| // 6. Verify cache cleared | |||
| assert.Equal(t, int64(-1), model.MatchEmailQuotaRule("user@full.com")) | |||
| } | |||
| func itoa(i int) string { | |||
| return strconv.Itoa(i) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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 | |||
| } | |||
| } | |||
| } | |||
| @@ -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") | |||
| } | |||
| @@ -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) | |||
| @@ -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) | |||
| } | |||
| @@ -61,12 +61,13 @@ func GetStatus(c *gin.Context) { | |||
| "telegram_oauth": common.TelegramOAuthEnabled, | |||
| "telegram_bot_name": common.TelegramBotName, | |||
| "system_name": common.SystemName, | |||
| "logo": common.Logo, | |||
| "logo": common.GetEffectiveLogo(), | |||
| "footer_html": common.Footer, | |||
| "wechat_qrcode": common.WeChatAccountQRCodeImageURL, | |||
| "wechat_login": common.WeChatAuthEnabled, | |||
| "server_address": system_setting.ServerAddress, | |||
| "turnstile_check": common.TurnstileCheckEnabled, | |||
| "captcha_enabled": common.CaptchaEnabled, | |||
| "turnstile_site_key": common.TurnstileSiteKey, | |||
| "top_up_link": common.TopUpLink, | |||
| "docs_link": operation_setting.GetGeneralSetting().DocsLink, | |||
| @@ -87,6 +88,7 @@ func GetStatus(c *gin.Context) { | |||
| "demo_site_enabled": operation_setting.DemoSiteEnabled, | |||
| "self_use_mode_enabled": operation_setting.SelfUseModeEnabled, | |||
| "default_use_auto_group": setting.DefaultUseAutoGroup, | |||
| "default_language": common.DefaultLanguage, | |||
| "usd_exchange_rate": operation_setting.USDExchangeRate, | |||
| "price": operation_setting.Price, | |||
| @@ -113,8 +115,10 @@ func GetStatus(c *gin.Context) { | |||
| "passkey_user_verification": passkeySetting.UserVerification, | |||
| "passkey_attachment": passkeySetting.AttachmentPreference, | |||
| "setup": constant.Setup, | |||
| "user_agreement_enabled": legalSetting.UserAgreement != "", | |||
| "privacy_policy_enabled": legalSetting.PrivacyPolicy != "", | |||
| "user_agreement_enabled": legalSetting.UserAgreementZh != "" || legalSetting.UserAgreementEn != "", | |||
| "privacy_policy_enabled": legalSetting.PrivacyPolicyZh != "" || legalSetting.PrivacyPolicyEn != "", | |||
| "terms_enabled": legalSetting.TermsOfServiceZh != "" || legalSetting.TermsOfServiceEn != "", | |||
| "usage_policy_enabled": legalSetting.UsagePolicyZh != "" || legalSetting.UsagePolicyEn != "", | |||
| "checkin_enabled": operation_setting.GetCheckinSetting().Enabled, | |||
| "_qn": "new-api", | |||
| } | |||
| @@ -188,20 +192,53 @@ func GetAbout(c *gin.Context) { | |||
| return | |||
| } | |||
| func getLegalContent(zh, en string, lang string) string { | |||
| if lang == "en" { | |||
| return en | |||
| } | |||
| return zh | |||
| } | |||
| func GetUserAgreement(c *gin.Context) { | |||
| ls := system_setting.GetLegalSettings() | |||
| lang := c.Query("lang") | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| "success": true, | |||
| "message": "", | |||
| "data": system_setting.GetLegalSettings().UserAgreement, | |||
| "data": getLegalContent(ls.UserAgreementZh, ls.UserAgreementEn, lang), | |||
| }) | |||
| return | |||
| } | |||
| func GetPrivacyPolicy(c *gin.Context) { | |||
| ls := system_setting.GetLegalSettings() | |||
| lang := c.Query("lang") | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| "success": true, | |||
| "message": "", | |||
| "data": getLegalContent(ls.PrivacyPolicyZh, ls.PrivacyPolicyEn, lang), | |||
| }) | |||
| return | |||
| } | |||
| func GetTermsOfService(c *gin.Context) { | |||
| ls := system_setting.GetLegalSettings() | |||
| lang := c.Query("lang") | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| "success": true, | |||
| "message": "", | |||
| "data": getLegalContent(ls.TermsOfServiceZh, ls.TermsOfServiceEn, lang), | |||
| }) | |||
| return | |||
| } | |||
| func GetUsagePolicy(c *gin.Context) { | |||
| ls := system_setting.GetLegalSettings() | |||
| lang := c.Query("lang") | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| "success": true, | |||
| "message": "", | |||
| "data": system_setting.GetLegalSettings().PrivacyPolicy, | |||
| "data": getLegalContent(ls.UsagePolicyZh, ls.UsagePolicyEn, lang), | |||
| }) | |||
| return | |||
| } | |||
| @@ -228,7 +265,39 @@ func GetHomePageContent(c *gin.Context) { | |||
| return | |||
| } | |||
| func GetCaptcha(c *gin.Context) { | |||
| if !common.CaptchaEnabled { | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| "success": false, | |||
| "message": "验证码功能未启用", | |||
| }) | |||
| return | |||
| } | |||
| id, b64s, err := common.GenerateCaptcha() | |||
| if err != nil { | |||
| common.ApiErrorMsg(c, "生成验证码失败") | |||
| return | |||
| } | |||
| common.ApiSuccess(c, gin.H{ | |||
| "id": id, | |||
| "captcha_image": b64s, | |||
| }) | |||
| } | |||
| func SendEmailVerification(c *gin.Context) { | |||
| if common.CaptchaEnabled { | |||
| captchaId := c.Query("captcha_id") | |||
| captchaCode := c.Query("captcha_code") | |||
| if captchaId == "" || captchaCode == "" { | |||
| common.ApiErrorMsg(c, "请先完成图片验证码") | |||
| return | |||
| } | |||
| if !common.VerifyCaptcha(captchaId, captchaCode) { | |||
| common.ApiErrorMsg(c, "图片验证码错误或已过期") | |||
| return | |||
| } | |||
| } | |||
| email := c.Query("email") | |||
| if err := common.Validate.Var(email, "required,email"); err != nil { | |||
| c.JSON(http.StatusOK, gin.H{ | |||
| @@ -0,0 +1,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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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 | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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}, | |||
| } | |||
| } | |||
| @@ -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, | |||
| }) | |||
| } | |||
| @@ -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) | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| } | |||
| @@ -87,6 +87,7 @@ func AddRedemption(c *gin.Context) { | |||
| cleanRedemption := model.Redemption{ | |||
| UserId: c.GetInt("id"), | |||
| Name: redemption.Name, | |||
| Remark: redemption.Remark, | |||
| Key: key, | |||
| CreatedTime: common.GetTimestamp(), | |||
| Quota: redemption.Quota, | |||
| @@ -146,6 +147,7 @@ func UpdateRedemption(c *gin.Context) { | |||
| } | |||
| // If you add more fields, please also update redemption.Update() | |||
| cleanRedemption.Name = redemption.Name | |||
| cleanRedemption.Remark = redemption.Remark | |||
| cleanRedemption.Quota = redemption.Quota | |||
| cleanRedemption.ExpiredTime = redemption.ExpiredTime | |||
| } | |||
| @@ -0,0 +1,159 @@ | |||
| package controller | |||
| import ( | |||
| "bytes" | |||
| "encoding/json" | |||
| "net/http" | |||
| "net/http/httptest" | |||
| "testing" | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/glebarez/sqlite" | |||
| "github.com/gin-gonic/gin" | |||
| "github.com/stretchr/testify/assert" | |||
| "github.com/stretchr/testify/require" | |||
| "gorm.io/gorm" | |||
| ) | |||
| func setupRedemptionControllerDB(t *testing.T) *gorm.DB { | |||
| t.Helper() | |||
| db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) | |||
| require.NoError(t, err) | |||
| sqlDB, _ := db.DB() | |||
| sqlDB.SetMaxOpenConns(1) | |||
| origDB := model.DB | |||
| origLogDB := model.LOG_DB | |||
| model.DB = db | |||
| model.LOG_DB = db | |||
| common.UsingSQLite = true | |||
| common.RedisEnabled = false | |||
| require.NoError(t, db.AutoMigrate(&model.User{}, &model.Redemption{}, &model.Log{})) | |||
| t.Cleanup(func() { | |||
| model.DB = origDB | |||
| model.LOG_DB = origLogDB | |||
| _ = sqlDB.Close() | |||
| }) | |||
| return db | |||
| } | |||
| func setupRedemptionRouter() *gin.Engine { | |||
| gin.SetMode(gin.TestMode) | |||
| r := gin.New() | |||
| r.Use(func(c *gin.Context) { | |||
| c.Set("id", 1) | |||
| c.Next() | |||
| }) | |||
| g := r.Group("/api/redemption") | |||
| g.POST("/", AddRedemption) | |||
| g.PUT("/", UpdateRedemption) | |||
| return r | |||
| } | |||
| func TestAddRedemptionStoresRemark(t *testing.T) { | |||
| db := setupRedemptionControllerDB(t) | |||
| router := setupRedemptionRouter() | |||
| body, err := json.Marshal(map[string]interface{}{ | |||
| "name": "campaign-a", | |||
| "remark": "admin only note", | |||
| "quota": 500000, | |||
| "count": 1, | |||
| "expired_time": 0, | |||
| }) | |||
| require.NoError(t, err) | |||
| req := httptest.NewRequest(http.MethodPost, "/api/redemption/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| w := httptest.NewRecorder() | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| var rows []model.Redemption | |||
| require.NoError(t, db.Find(&rows).Error) | |||
| require.Len(t, rows, 1) | |||
| assert.Equal(t, "campaign-a", rows[0].Name) | |||
| assert.Equal(t, "admin only note", rows[0].Remark) | |||
| } | |||
| func TestUpdateRedemptionStoresRemark(t *testing.T) { | |||
| db := setupRedemptionControllerDB(t) | |||
| router := setupRedemptionRouter() | |||
| row := model.Redemption{ | |||
| Id: 1, | |||
| UserId: 1, | |||
| Key: "update-remark-key", | |||
| Name: "campaign-b", | |||
| Remark: "before update", | |||
| Status: common.RedemptionCodeStatusEnabled, | |||
| Quota: 500000, | |||
| CreatedTime: common.GetTimestamp(), | |||
| } | |||
| require.NoError(t, db.Create(&row).Error) | |||
| body, err := json.Marshal(map[string]interface{}{ | |||
| "id": 1, | |||
| "name": "campaign-b", | |||
| "remark": "after update", | |||
| "quota": 500000, | |||
| "expired_time": 0, | |||
| }) | |||
| require.NoError(t, err) | |||
| req := httptest.NewRequest(http.MethodPut, "/api/redemption/", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| w := httptest.NewRecorder() | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| var stored model.Redemption | |||
| require.NoError(t, db.First(&stored, 1).Error) | |||
| assert.Equal(t, "after update", stored.Remark) | |||
| } | |||
| func TestUpdateRedemptionStatusOnlyKeepsRemark(t *testing.T) { | |||
| db := setupRedemptionControllerDB(t) | |||
| router := setupRedemptionRouter() | |||
| row := model.Redemption{ | |||
| Id: 1, | |||
| UserId: 1, | |||
| Key: "status-only-remark-key", | |||
| Name: "campaign-c", | |||
| Remark: "keep this remark", | |||
| Status: common.RedemptionCodeStatusEnabled, | |||
| Quota: 500000, | |||
| CreatedTime: common.GetTimestamp(), | |||
| } | |||
| require.NoError(t, db.Create(&row).Error) | |||
| body, err := json.Marshal(map[string]interface{}{ | |||
| "id": 1, | |||
| "status": common.RedemptionCodeStatusDisabled, | |||
| "remark": "should not overwrite", | |||
| }) | |||
| require.NoError(t, err) | |||
| req := httptest.NewRequest(http.MethodPut, "/api/redemption/?status_only=true", bytes.NewReader(body)) | |||
| req.Header.Set("Content-Type", "application/json") | |||
| w := httptest.NewRecorder() | |||
| router.ServeHTTP(w, req) | |||
| assert.Equal(t, http.StatusOK, w.Code) | |||
| var stored model.Redemption | |||
| require.NoError(t, db.First(&stored, 1).Error) | |||
| assert.Equal(t, common.RedemptionCodeStatusDisabled, stored.Status) | |||
| assert.Equal(t, "keep this remark", stored.Remark) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -0,0 +1,55 @@ | |||
| package controller | |||
| import ( | |||
| "github.com/QuantumNous/new-api/common" | |||
| "github.com/QuantumNous/new-api/model" | |||
| "github.com/gin-gonic/gin" | |||
| ) | |||
| // ReorderVendors 批量更新供应商排序 | |||
| func ReorderVendors(c *gin.Context) { | |||
| var req struct { | |||
| Items []struct { | |||
| Id int `json:"id"` | |||
| SortOrder int `json:"sort_order"` | |||
| } `json:"items"` | |||
| } | |||
| if err := c.ShouldBindJSON(&req); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| if len(req.Items) == 0 { | |||
| common.ApiErrorMsg(c, "items 不能为空") | |||
| return | |||
| } | |||
| if err := model.ReorderVendors(req.Items); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| common.ApiSuccess(c, nil) | |||
| } | |||
| // ReorderModels 批量更新模型排序 | |||
| func ReorderModels(c *gin.Context) { | |||
| var req struct { | |||
| Items []struct { | |||
| Id int `json:"id"` | |||
| SortOrder int `json:"sort_order"` | |||
| } `json:"items"` | |||
| } | |||
| if err := c.ShouldBindJSON(&req); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| if len(req.Items) == 0 { | |||
| common.ApiErrorMsg(c, "items 不能为空") | |||
| return | |||
| } | |||
| if err := model.ReorderModels(req.Items); err != nil { | |||
| common.ApiError(c, err) | |||
| return | |||
| } | |||
| model.RefreshPricing() | |||
| common.ApiSuccess(c, nil) | |||
| } | |||
| @@ -72,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) | |||
| @@ -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, | |||
| })) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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 | |||
| @@ -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) | |||
| } | |||
| @@ -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"]) | |||
| } | |||
| @@ -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": "", | |||
| }) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -0,0 +1,106 @@ | |||
| package controller | |||
| import ( | |||
| "bytes" | |||
| "net/http" | |||