diff --git a/.github/audit-exceptions.yml b/.github/audit-exceptions.yml index 4e05aae66b..2a89dd0761 100644 --- a/.github/audit-exceptions.yml +++ b/.github/audit-exceptions.yml @@ -5,14 +5,14 @@ exceptions: severity: high reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2023-30533)" mitigation: "Load only on export; restrict export permissions and data scope" - expires_on: "2026-07-06" + expires_on: "2026-10-06" owner: "security@your-domain" - package: xlsx advisory: "GHSA-5pgg-2g8v-p4x9" severity: high reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2024-22363)" mitigation: "Load only on export; restrict export permissions and data scope" - expires_on: "2026-07-06" + expires_on: "2026-10-06" owner: "security@your-domain" - package: lodash advisory: "GHSA-r5fr-rjxr-66jc" diff --git a/README.md b/README.md index d960d93fa1..eb24b0d3d6 100644 --- a/README.md +++ b/README.md @@ -23,10 +23,11 @@ Please read the following carefully before using this project: - **🚨 Terms of Service Risk**: Using this project may violate the terms of service of Anthropic and other upstream providers. Please review the relevant providers' user agreements before use; all risks arising from such use are borne solely by the user. - **⚖️ Compliant Use**: Use this project only in compliance with the laws and regulations of your country or region. Any unlawful use is strictly prohibited. - **📖 Disclaimer**: This project is provided for technical learning and research purposes only. The authors assume no liability for account bans, service interruptions, data loss, or any other direct or indirect damages resulting from the use of this project. +- **🚫 No Commercial Authorization**: The developers of this project have never authorized any individual or organization to conduct any form of commercial operation based on this project. Any commercial activity conducted in the name of or based on this project is unrelated to this project and its developers, and all resulting disputes, losses, and legal liabilities shall be borne solely by the party conducting such activity. ## ❤️ Sponsors -> [Want to appear here?](mailto:support@pincc.ai) +> [Want to appear here?](mailto:support@sub2api.org) @@ -140,7 +141,7 @@ Model authenticity: no content intervention or secondary filtering — experienc - @@ -522,20 +523,20 @@ Additional security-related options are available in `config.yaml`: **⚠️ Security Warning: HTTP URL Configuration** -When `security.url_allowlist.enabled=false`, the system performs minimal URL validation by default, **rejecting HTTP URLs** and only allowing HTTPS. To allow HTTP URLs (e.g., for development or internal testing), you must explicitly set: +When `security.url_allowlist.enabled=false`, the system performs minimal URL validation and **allows HTTP URLs by default** (dev-friendly mode; Docker Compose deployments use the same default). For production, explicitly tighten this to HTTPS-only: ```yaml security: url_allowlist: enabled: false # Disable allowlist checks - allow_insecure_http: true # Allow HTTP URLs (⚠️ INSECURE) + allow_insecure_http: false # HTTPS only (recommended for production) ``` **Or via environment variable:** ```bash SECURITY_URL_ALLOWLIST_ENABLED=false -SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true +SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=false ``` **Risks of allowing HTTP:** @@ -549,7 +550,7 @@ SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true - ✅ Testing account connectivity before obtaining HTTPS - ❌ Production environments (use HTTPS only) -**Example error without this setting:** +**Example error for HTTP URLs when `allow_insecure_http: false` is set:** ``` Invalid base URL: invalid url scheme: http ``` @@ -630,8 +631,10 @@ Sub2API supports Grok subscription accounts through xAI OAuth and forwards OpenA - Public Claude-compatible target: `/v1/messages`, converted to xAI Responses and returned as Anthropic Messages output for Claude CLI style clients - Public Chat Completions targets: `/v1/chat/completions` and `/chat/completions`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/chat/completions` - Codex CLI style Responses WebSocket ingress is accepted on the Responses targets and bridged to xAI HTTP/SSE Responses upstream -- Initial models: `grok-4.3`, `grok-build-0.1`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309` -- Out of scope for this provider: image, video, TTS, transcription, browser automation, cookies, and Grok web scraping +- Initial text models: `grok-4.3`, `grok-build-0.1`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309` +- Media targets for Grok groups: `/v1/images/generations`, `/images/generations`, `/v1/images/edits`, `/images/edits`, `/v1/videos/generations`, `/videos/generations`, `/v1/videos/{request_id}`, and `/videos/{request_id}`. Generation requests require the group image-generation permission. +- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5` +- Out of scope for this provider: TTS, transcription, browser automation, cookies, and Grok web scraping ### OAuth Configuration @@ -689,12 +692,6 @@ Antigravity accounts support optional **hybrid scheduling**. When enabled, the g > **⚠️ Warning**: Anthropic Claude and Antigravity Claude **cannot be mixed within the same conversation context**. Use groups to isolate them properly. -### Known Issues - -In Claude Code, Plan Mode cannot exit automatically. (Normally when using the native Claude API, after planning is complete, Claude Code will pop up options for users to approve or reject the plan.) - -**Workaround**: Press `Shift + Tab` to manually exit Plan Mode, then type your response to approve or reject the plan. - --- ## Project Structure @@ -725,16 +722,6 @@ sub2api/ └── install.sh # One-click installation script ``` -## Disclaimer - -> **Please read carefully before using this project:** -> -> :rotating_light: **Terms of Service Risk**: Using this project may violate Anthropic's Terms of Service. Please read Anthropic's user agreement carefully before use. All risks arising from the use of this project are borne solely by the user. -> -> :book: **Disclaimer**: This project is for technical learning and research purposes only. The author assumes no responsibility for account suspension, service interruption, or any other losses caused by the use of this project. - ---- - ## Star History diff --git a/README_CN.md b/README_CN.md index 2ba06a5299..a93056282b 100644 --- a/README_CN.md +++ b/README_CN.md @@ -24,10 +24,11 @@ - **🚨 服务条款风险**:使用本项目可能违反 Anthropic 等上游服务商的服务条款。请在使用前仔细阅读相关服务商的用户协议,由此产生的一切风险由用户自行承担。 - **⚖️ 合规使用**:请在符合您所在国家或地区法律法规的前提下使用本项目,严禁将其用于任何违法违规用途。 - **📖 免责声明**:本项目仅供技术学习与研究使用,作者不对因使用本项目导致的账户封禁、服务中断、数据丢失或其他任何直接或间接损失承担责任。 +- **🚫 无商业授权**:本项目从未授权任何个人或组织基于本项目开展任何形式的商业化运营。任何以本项目名义或基于本项目从事的商业行为均与本项目及其开发者无关,由此产生的一切纠纷、损失和法律责任由行为主体自行承担。 ## ❤️ 赞助商 -> [想出现在这里?](mailto:support@pincc.ai) +> [想出现在这里?](mailto:support@sub2api.org)
proxy4freeProxy4Free is a data proxy service provider for developers and AI applications, offering residential proxies, static residential proxies, ISP proxies, and datacenter proxies for scenarios such as Web Scraping, Browser Automation, and AI Agents. With global IP resources, stable connections, and flexible switching, it helps developers improve data collection success rates and reduce the risk of IP bans. Register via this link to get started and easily build more stable and efficient automation workflows. +Thanks to Proxy4Free for sponsoring this project! Proxy4Free is a data proxy service provider for developers and AI applications, offering residential proxies, static residential proxies, ISP proxies, and datacenter proxies for scenarios such as Web Scraping, Browser Automation, and AI Agents. With global IP resources, stable connections, and flexible switching, it helps developers improve data collection success rates and reduce the risk of IP bans. Register via this link to get started and easily build more stable and efficient automation workflows.
@@ -143,10 +144,11 @@ - +
proxy4freeProxy4Free 是面向开发者和 AI 应用的数据代理服务商,提供住宅代理、静态住宅代理、ISP 代理及数据中心代理等多种代理解决方案,适用于 Web Scraping、Browser Automation、AI Agent 等场景。支持全球 IP 资源、稳定连接与灵活切换,帮助开发者提升数据采集成功率,降低 IP 封禁风险。通过此链接注册即可开始体验,轻松构建更稳定、高效的自动化工作流。 +感谢 Proxy4Free 赞助本项目!Proxy4Free 是面向开发者和 AI 应用的数据代理服务商,提供住宅代理、静态住宅代理、ISP 代理及数据中心代理等多种代理解决方案,适用于 Web Scraping、Browser Automation、AI Agent 等场景。支持全球 IP 资源、稳定连接与灵活切换,帮助开发者提升数据采集成功率,降低 IP 封禁风险。通过此链接注册即可开始体验,轻松构建更稳定、高效的自动化工作流。
## 项目概述 @@ -566,20 +568,20 @@ gateway: **⚠️ 安全警告:HTTP URL 配置** -当 `security.url_allowlist.enabled=false` 时,系统默认执行最小 URL 校验,**拒绝 HTTP URL**,仅允许 HTTPS。要允许 HTTP URL(例如用于开发或内网测试),必须显式设置: +当 `security.url_allowlist.enabled=false` 时,系统仅执行最小 URL 校验,且**默认允许 HTTP URL**(开发友好模式,Docker Compose 部署的默认值一致)。生产环境建议显式收紧为仅允许 HTTPS: ```yaml security: url_allowlist: enabled: false # 禁用白名单检查 - allow_insecure_http: true # 允许 HTTP URL(⚠️ 不安全) + allow_insecure_http: false # 仅允许 HTTPS(生产环境推荐) ``` **或通过环境变量:** ```bash SECURITY_URL_ALLOWLIST_ENABLED=false -SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true +SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=false ``` **允许 HTTP 的风险:** @@ -593,7 +595,7 @@ SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true - ✅ 获取 HTTPS 前测试账号连通性 - ❌ 生产环境(仅使用 HTTPS) -**未设置此项时的错误示例:** +**设置 `allow_insecure_http: false` 后,HTTP URL 会返回如下错误:** ``` Invalid base URL: invalid url scheme: http ``` @@ -709,10 +711,6 @@ Antigravity 账户支持可选的**混合调度**功能。开启后,通用端 > **⚠️ 注意**:Anthropic Claude 和 Antigravity Claude **不能在同一上下文中混合使用**,请通过分组功能做好隔离。 - -### 已知问题 -在 Claude Code 中,无法自动退出Plan Mode。(正常使用原生Claude Api时,Plan 完成后,Claude Code会弹出弹出选项让用户同意或拒绝Plan。) -解决办法:shift + Tab,手动退出Plan mode,然后输入内容 告诉 Claude Code 同意或拒绝 Plan --- ## 项目结构 @@ -743,16 +741,6 @@ sub2api/ └── install.sh # 一键安装脚本 ``` -## 免责声明 - -> **使用本项目前请仔细阅读:** -> -> :rotating_light: **服务条款风险**: 使用本项目可能违反 Anthropic 的服务条款。请在使用前仔细阅读 Anthropic 的用户协议,使用本项目的一切风险由用户自行承担。 -> -> :book: **免责声明**: 本项目仅供技术学习和研究使用,作者不对因使用本项目导致的账户封禁、服务中断或其他损失承担任何责任。 - ---- - ## Star History diff --git a/README_JA.md b/README_JA.md index 3772b93957..bd154a9a6b 100644 --- a/README_JA.md +++ b/README_JA.md @@ -23,10 +23,11 @@ - **🚨 利用規約のリスク**:本プロジェクトの使用は、Anthropic をはじめとする上流プロバイダーの利用規約に違反する可能性があります。ご利用前に各プロバイダーのユーザー規約を必ずご確認ください。使用により生じるすべてのリスクはユーザーご自身が負うものとします。 - **⚖️ 法令遵守**:お住まいの国または地域の法令を遵守した上で本プロジェクトをご利用ください。いかなる違法な目的での使用も固く禁じます。 - **📖 免責事項**:本プロジェクトは技術的な学習および研究の目的でのみ提供されます。本プロジェクトの使用により生じたアカウントの停止、サービスの中断、データの損失、その他一切の直接的または間接的な損害について、作者は一切の責任を負いません。 +- **🚫 商用利用の非許諾**:本プロジェクトの開発者は、いかなる個人または組織に対しても、本プロジェクトを利用したいかなる形態の商業運営も一切許諾していません。本プロジェクトの名義で、または本プロジェクトに基づいて行われる商業行為はすべて本プロジェクトおよびその開発者とは無関係であり、それにより生じる一切の紛争、損失、法的責任は行為者自身が負うものとします。 ## ❤️ スポンサー -> [こちらに掲載しませんか?](mailto:support@pincc.ai) +> [こちらに掲載しませんか?](mailto:support@sub2api.org) @@ -138,7 +139,7 @@ - @@ -520,20 +521,20 @@ default: **⚠️ セキュリティ警告: HTTP URL 設定** -`security.url_allowlist.enabled=false` の場合、システムはデフォルトで最小限の URL バリデーションを行い、**HTTP URL を拒否**して HTTPS のみを許可します。HTTP URL を許可するには(開発環境や内部テスト用など)、以下を明示的に設定する必要があります: +`security.url_allowlist.enabled=false` の場合、システムは最小限の URL バリデーションのみを行い、**デフォルトで HTTP URL を許可**します(開発フレンドリーモード。Docker Compose デプロイのデフォルトも同じです)。本番環境では、以下のように明示的に HTTPS のみに制限することを推奨します: ```yaml security: url_allowlist: enabled: false # 許可リストチェックを無効化 - allow_insecure_http: true # HTTP URL を許可(⚠️ セキュリティリスクあり) + allow_insecure_http: false # HTTPS のみ許可(本番環境推奨) ``` **または環境変数で設定:** ```bash SECURITY_URL_ALLOWLIST_ENABLED=false -SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true +SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=false ``` **HTTP を許可するリスク:** @@ -547,7 +548,7 @@ SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true - ✅ HTTPS 取得前のアカウント接続テスト - ❌ 本番環境(HTTPS のみを使用) -**この設定なしで表示されるエラー例:** +**`allow_insecure_http: false` 設定時に HTTP URL で表示されるエラー例:** ``` Invalid base URL: invalid url scheme: http ``` @@ -640,12 +641,6 @@ Antigravity アカウントはオプションの**ハイブリッドスケジュ > **⚠️ 警告**: Anthropic Claude と Antigravity Claude は**同じ会話コンテキスト内で混在させることはできません**。グループを使用して適切に分離してください。 -### 既知の問題 - -Claude Code では、Plan Mode を自動的に終了できません。(通常、ネイティブの Claude API を使用する場合、計画が完了すると Claude Code はユーザーに計画を承認または拒否するオプションをポップアップ表示します。) - -**回避策**: `Shift + Tab` を押して手動で Plan Mode を終了し、計画を承認または拒否するためのレスポンスを入力してください。 - --- ## プロジェクト構成 @@ -676,16 +671,6 @@ sub2api/ └── install.sh # ワンクリックインストールスクリプト ``` -## 免責事項 - -> **本プロジェクトをご利用の前に、以下をよくお読みください:** -> -> :rotating_light: **利用規約違反のリスク**: 本プロジェクトの使用は Anthropic の利用規約に違反する可能性があります。使用前に Anthropic のユーザー契約をよくお読みください。本プロジェクトの使用に起因するすべてのリスクは、ユーザー自身が負うものとします。 -> -> :book: **免責事項**: 本プロジェクトは技術的な学習および研究目的のみで提供されています。作者は、本プロジェクトの使用によるアカウント停止、サービス中断、その他の損失について一切の責任を負いません。 - ---- - ## スター履歴 diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 3170382e8b..a0e8ec1d4e 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.143 +0.1.145 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index a6fb5266aa..a412563a6c 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -67,7 +67,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { serviceUserPlatformQuotaRepository := repository.NewUserPlatformQuotaServiceAdapter(userPlatformQuotaRepository) billingCacheService := service.ProvideBillingCacheService(billingCache, userRepository, userSubscriptionRepository, apiKeyRepository, userRPMCache, userGroupRateRepository, configConfig, serviceUserPlatformQuotaRepository) apiKeyCache := repository.NewAPIKeyCache(redisClient) - apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService) + concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig) + schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig) + accountRepository := repository.NewAccountRepository(client, db, schedulerCache) + concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig) + apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService, concurrencyService) apiKeyAuthCacheInvalidator := service.ProvideAPIKeyAuthCacheInvalidator(apiKeyService) promoService := service.NewPromoService(promoCodeRepository, userRepository, billingCacheService, client, apiKeyAuthCacheInvalidator) subscriptionService := service.NewSubscriptionService(groupRepository, userSubscriptionRepository, billingCacheService, client, configConfig) @@ -92,10 +96,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { usageLogRepository := repository.NewUsageLogRepository(client, db) usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator) opsRepository := repository.NewOpsRepository(db) - schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig) - accountRepository := repository.NewAccountRepository(client, db, schedulerCache) - concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig) - concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig) usageBillingRepository := repository.NewUsageBillingRepository(client, db) gatewayCache := repository.NewGatewayCache(redisClient) schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db) diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 18baa34881..99fedb5b1c 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -969,6 +969,9 @@ type GatewayOpenAIWSSchedulerScoreWeights struct { Reset float64 `mapstructure:"reset"` // QuotaHeadroom 倾向 7d 剩余额度更健康的账号;默认 0(关闭,不改变原有行为)。 QuotaHeadroom float64 `mapstructure:"quota_headroom"` + // PreviousResponse/SessionSticky 仅在开启 OpenAI 高级调度的粘性加权时生效。 + PreviousResponse float64 `mapstructure:"previous_response"` + SessionSticky float64 `mapstructure:"session_sticky"` } // GatewayOpenAISchedulerConfig OpenAI 高级调度器配置。 @@ -1891,6 +1894,8 @@ func setDefaults() { viper.SetDefault("gateway.openai_ws.scheduler_score_weights.ttft", 0.5) viper.SetDefault("gateway.openai_ws.scheduler_score_weights.reset", 0.0) viper.SetDefault("gateway.openai_ws.scheduler_score_weights.quota_headroom", 0.0) + viper.SetDefault("gateway.openai_ws.scheduler_score_weights.previous_response", 5.0) + viper.SetDefault("gateway.openai_ws.scheduler_score_weights.session_sticky", 3.0) // OpenAI HTTP upstream protocol strategy viper.SetDefault("gateway.openai_http2.enabled", true) viper.SetDefault("gateway.openai_http2.allow_proxy_fallback_to_http1", true) @@ -1945,7 +1950,10 @@ func setDefaults() { viper.SetDefault("gateway.usage_record.worker_count", 128) viper.SetDefault("gateway.usage_record.queue_size", 16384) viper.SetDefault("gateway.usage_record.task_timeout_seconds", 5) - viper.SetDefault("gateway.usage_record.overflow_policy", UsageRecordOverflowPolicySample) + // 默认 sync:队列满时由提交方内联执行(提交点在响应写出之后,不阻塞客户端)。 + // sample/drop 会在溢出时静默丢弃计费任务,造成扣费与 usage_logs 对账缺口(issue #3656), + // 仅供显式配置的运维场景使用。 + viper.SetDefault("gateway.usage_record.overflow_policy", UsageRecordOverflowPolicySync) viper.SetDefault("gateway.usage_record.overflow_sample_percent", 10) viper.SetDefault("gateway.usage_record.auto_scale_enabled", true) viper.SetDefault("gateway.usage_record.auto_scale_min_workers", 128) @@ -2671,7 +2679,9 @@ func (c *Config) Validate() error { c.Gateway.OpenAIWS.SchedulerScoreWeights.Queue < 0 || c.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate < 0 || c.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT < 0 || - c.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom < 0 { + c.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom < 0 || + c.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse < 0 || + c.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky < 0 { return fmt.Errorf("gateway.openai_ws.scheduler_score_weights.* must be non-negative") } weightSum := c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority + diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index bf7a327563..804155d1a9 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -1903,8 +1903,8 @@ func TestLoad_DefaultGatewayUsageRecordConfig(t *testing.T) { if cfg.Gateway.UsageRecord.TaskTimeoutSeconds != 5 { t.Fatalf("task_timeout_seconds = %d, want 5", cfg.Gateway.UsageRecord.TaskTimeoutSeconds) } - if cfg.Gateway.UsageRecord.OverflowPolicy != UsageRecordOverflowPolicySample { - t.Fatalf("overflow_policy = %s, want %s", cfg.Gateway.UsageRecord.OverflowPolicy, UsageRecordOverflowPolicySample) + if cfg.Gateway.UsageRecord.OverflowPolicy != UsageRecordOverflowPolicySync { + t.Fatalf("overflow_policy = %s, want %s", cfg.Gateway.UsageRecord.OverflowPolicy, UsageRecordOverflowPolicySync) } if cfg.Gateway.UsageRecord.OverflowSamplePercent != 10 { t.Fatalf("overflow_sample_percent = %d, want 10", cfg.Gateway.UsageRecord.OverflowSamplePercent) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 0d3bf88705..044e07d7a2 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -68,6 +68,9 @@ const ( SubscriptionStatusSuspended = "suspended" ) +// AntigravityGemini31ProAgentModel is the upstream route for Gemini 3.1 Pro High. +const AntigravityGemini31ProAgentModel = "gemini-pro-agent" + // DefaultAntigravityModelMapping 是 Antigravity 平台的默认模型映射 // 当账号未配置 model_mapping 时使用此默认值 // 与前端 useModelWhitelist.ts 中的 antigravityDefaultMappings 保持一致 @@ -103,10 +106,12 @@ var DefaultAntigravityModelMapping = map[string]string{ "gemini-3-flash-preview": "gemini-3-flash", "gemini-3-pro-preview": "gemini-3-pro-high", // Gemini 3.1 白名单 - "gemini-3.1-pro-high": "gemini-3.1-pro-high", - "gemini-3.1-pro-low": "gemini-3.1-pro-low", + AntigravityGemini31ProAgentModel: AntigravityGemini31ProAgentModel, + "gemini-3.1-pro": AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-high": AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-low": "gemini-3.1-pro-low", // Gemini 3.1 preview 映射 - "gemini-3.1-pro-preview": "gemini-3.1-pro-high", + "gemini-3.1-pro-preview": AntigravityGemini31ProAgentModel, // Gemini 3.1 image 白名单 "gemini-3.1-flash-image": "gemini-3.1-flash-image", // Gemini 3.1 image preview 映射 diff --git a/backend/internal/domain/constants_test.go b/backend/internal/domain/constants_test.go index 0b24aea915..0fb9054f7e 100644 --- a/backend/internal/domain/constants_test.go +++ b/backend/internal/domain/constants_test.go @@ -43,6 +43,28 @@ func TestDefaultAntigravityModelMapping_ContainsNewClaudeModels(t *testing.T) { } } +func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) { + t.Parallel() + + cases := map[string]string{ + AntigravityGemini31ProAgentModel: AntigravityGemini31ProAgentModel, + "gemini-3.1-pro": AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-high": AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-preview": AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-low": "gemini-3.1-pro-low", + } + + for from, want := range cases { + got, ok := DefaultAntigravityModelMapping[from] + if !ok { + t.Fatalf("expected mapping for %q to exist", from) + } + if got != want { + t.Fatalf("unexpected mapping for %q: got %q want %q", from, got, want) + } + } +} + func TestDefaultBedrockModelMapping_ContainsNewClaudeModels(t *testing.T) { t.Parallel() diff --git a/backend/internal/handler/admin/account_codex_import.go b/backend/internal/handler/admin/account_codex_import.go index 6ec9495d84..01a5fbfa1c 100644 --- a/backend/internal/handler/admin/account_codex_import.go +++ b/backend/internal/handler/admin/account_codex_import.go @@ -253,6 +253,17 @@ func (h *AccountHandler) importCodexSessions(ctx context.Context, req CodexSessi Message: "已有账号未记录 chatgpt_user_id,已按共享的 chatgpt_account_id 匹配并回填,请确认两者属于同一用户", }) } + preserveExistingRefresh := item.RefreshToken == "" && + codexCredentialString(existing.Credentials, "refresh_token") != "" + if preserveExistingRefresh { + result.Warnings = append(result.Warnings, CodexSessionImportMessage{ + Index: entry.Index, + Name: accountName, + Message: "已有账号包含 refresh_token,本次 accessToken-only 导入已保留自动续期凭据", + }) + effectiveExpiresAt = nil + autoPauseOnExpired = nil + } mergedCredentials := mergeCodexImportCredentials(existing.Credentials, credentials, item) mergedExtra := mergeCodexImportMap(existing.Extra, extra) updateInput := &service.UpdateAccountInput{ @@ -592,7 +603,7 @@ func normalizeCodexImportEntry(entry codexImportEntry) (*codexImportAccount, err fingerprint := codexTokenFingerprint(item.AccessToken) item.Extra["access_token_sha256"] = fingerprint - item.IdentityKeys = buildCodexIdentityKeys(item.AccountID, item.UserID, item.Email, item.AccessToken) + item.IdentityKeys = buildCodexImportIdentityKeys(item.AccountID, item.UserID, item.Email, item.AccessToken, item.RefreshToken) item.Name = buildCodexImportAccountName(item, entry.Index) return item, nil @@ -815,13 +826,25 @@ func sanitizeCodexImportCredentialExtras(input map[string]any) map[string]any { return out } -// buildCodexIdentityKeys 按身份强度排序生成匹配键:chatgpt_account_id 在同一 -// ChatGPT 团队内是共享的,因此 account: 键排在最后,且命中时还需通过 -// codexIdentityConflicts 的跨用户校验才生效。 -func buildCodexIdentityKeys(accountID, userID, email, accessToken string) []string { +// buildCodexImportIdentityKeys 生成导入条目的匹配键。refresh_token 缺失时 +// Codex session 只能作为 accessToken-only 凭据使用,此时以 access token +// 指纹作为唯一稳定身份,避免同 workspace 下共享的 account/user 标识误合并。 +func buildCodexImportIdentityKeys(accountID, userID, email, accessToken, refreshToken string) []string { + accessToken = strings.TrimSpace(accessToken) + refreshToken = strings.TrimSpace(refreshToken) + if refreshToken == "" && accessToken != "" { + return []string{"access:" + codexTokenFingerprint(accessToken)} + } + return buildCodexStoredIdentityKeys(accountID, userID, email, accessToken) +} + +// buildCodexStoredIdentityKeys 生成存量账号索引键,保留 user/account 维度, +// 让 accessToken-only 账号后续升级为完整 OAuth 时仍能命中并更新原账号。 +func buildCodexStoredIdentityKeys(accountID, userID, email, accessToken string) []string { keys := make([]string, 0, 3) accountID = strings.TrimSpace(accountID) userID = strings.TrimSpace(userID) + accessToken = strings.TrimSpace(accessToken) if userID != "" { keys = append(keys, "user:"+userID) } @@ -830,7 +853,7 @@ func buildCodexIdentityKeys(accountID, userID, email, accessToken string) []stri keys = append(keys, "email:"+email) } } - if accessToken = strings.TrimSpace(accessToken); accessToken != "" { + if accessToken != "" { keys = append(keys, "access:"+codexTokenFingerprint(accessToken)) } if accountID != "" { @@ -854,7 +877,8 @@ func (i *codexAccountIndex) Add(account service.Account) { if i.accountsByKey == nil { i.accountsByKey = map[string][]service.Account{} } - keys := buildCodexIdentityKeys( + i.remove(account.ID) + keys := buildCodexStoredIdentityKeys( codexCredentialString(account.Credentials, "chatgpt_account_id"), codexCredentialString(account.Credentials, "chatgpt_user_id"), codexCredentialString(account.Credentials, "email"), @@ -865,6 +889,22 @@ func (i *codexAccountIndex) Add(account service.Account) { } } +func (i *codexAccountIndex) remove(accountID int64) { + for key, accounts := range i.accountsByKey { + kept := accounts[:0] + for _, account := range accounts { + if account.ID != accountID { + kept = append(kept, account) + } + } + if len(kept) == 0 { + delete(i.accountsByKey, key) + continue + } + i.accountsByKey[key] = kept + } +} + // upsertCodexAccount 保留同一键下的全部候选账号(共享的 account: 键可对应 // 团队内多个账号),同一账号重复 Add 时原位替换为最新状态。 func upsertCodexAccount(accounts []service.Account, account service.Account) []service.Account { @@ -894,9 +934,9 @@ func (i *codexAccountIndex) Find(keys []string, userID string) (*service.Account } // codexIdentityConflicts 判断 account: 键的命中是否把同一 ChatGPT 团队的两个 -// 不同成员误连到一起:双方都携带 user id 且不相等时视为冲突。任一侧缺少 -// user id 时保留匹配,使早期未记录 chatgpt_user_id 的存量账号仍能被更新 -// (并借助凭据合并回填 user id),而不是产生重复账号。 +// 不同成员误连到一起:双方都携带 user id 且不相等时视为冲突。存量索引侧 +// 仍保留 account 键,任一侧缺少 user id 时允许匹配,使含 refresh_token +// 的常规导入和 accessToken-only 账号升级为完整 OAuth 时仍能更新原账号。 func codexIdentityConflicts(key, userID, storedUserID string) bool { if !strings.HasPrefix(key, "account:") { return false @@ -948,8 +988,15 @@ func mergeCodexImportCredentials(existing, incoming map[string]any, item *codexI return out } if strings.TrimSpace(item.RefreshToken) == "" { - delete(out, "refresh_token") - delete(out, "client_id") + if codexCredentialString(existing, "refresh_token") == "" { + delete(out, "refresh_token") + delete(out, "client_id") + } else { + out["refresh_token"] = existing["refresh_token"] + if clientID, ok := existing["client_id"]; ok { + out["client_id"] = clientID + } + } } if strings.TrimSpace(item.IDToken) == "" { delete(out, "id_token") diff --git a/backend/internal/handler/admin/account_codex_import_test.go b/backend/internal/handler/admin/account_codex_import_test.go index f4ee5bd7db..a52463aa86 100644 --- a/backend/internal/handler/admin/account_codex_import_test.go +++ b/backend/internal/handler/admin/account_codex_import_test.go @@ -1,6 +1,7 @@ package admin import ( + "context" "encoding/base64" "encoding/json" "fmt" @@ -144,7 +145,7 @@ func TestNormalizeCodexSessionJSONExtractsCredentialsAndIgnoresSessionToken(t *t } } -func TestMergeCodexImportCredentialsClearsStaleRefreshFieldsWhenIncomingHasNoRefreshToken(t *testing.T) { +func TestMergeCodexImportCredentialsPreservesExistingRefreshFieldsWhenIncomingHasNoRefreshToken(t *testing.T) { existing := map[string]any{ "access_token": "old-access-token", "refresh_token": "old-refresh-token", @@ -171,11 +172,11 @@ func TestMergeCodexImportCredentialsClearsStaleRefreshFieldsWhenIncomingHasNoRef if merged["chatgpt_account_id"] != "acct-new" { t.Fatalf("chatgpt_account_id = %v, want acct-new", merged["chatgpt_account_id"]) } - if _, ok := merged["refresh_token"]; ok { - t.Fatalf("refresh_token should be cleared") + if merged["refresh_token"] != "old-refresh-token" { + t.Fatalf("refresh_token = %v, want old-refresh-token", merged["refresh_token"]) } - if _, ok := merged["client_id"]; ok { - t.Fatalf("client_id should be cleared") + if merged["client_id"] != "old-client-id" { + t.Fatalf("client_id = %v, want old-client-id", merged["client_id"]) } if _, ok := merged["id_token"]; ok { t.Fatalf("id_token should be cleared") @@ -301,9 +302,9 @@ func TestResolveCodexImportExpiryForNoRefreshTokenUsesEarlierRequestExpiry(t *te } func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) { - keys := buildCodexIdentityKeys("acct-1", "user-1", "same@example.com", "token") + keys := buildCodexImportIdentityKeys("acct-1", "user-1", "same@example.com", "token", "refresh") if len(keys) == 0 || keys[0] != "user:user-1" { - t.Fatalf("user key should have highest priority: %v", keys) + t.Fatalf("user key should have highest priority when refresh token exists: %v", keys) } if keys[len(keys)-1] != "account:acct-1" { t.Fatalf("shared account key should be the last fallback: %v", keys) @@ -314,7 +315,7 @@ func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) { } } - keys = buildCodexIdentityKeys("", "", "same@example.com", "token") + keys = buildCodexImportIdentityKeys("", "", "same@example.com", "token", "refresh") hasEmail := false for _, key := range keys { if key == "email:same@example.com" { @@ -324,6 +325,11 @@ func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) { if !hasEmail { t.Fatalf("weak identity should include email fallback: %v", keys) } + + keys = buildCodexImportIdentityKeys("acct-1", "user-1", "same@example.com", "token", "") + if len(keys) != 1 || !strings.HasPrefix(keys[0], "access:") { + t.Fatalf("accessToken-only identity should use only access fingerprint: %v", keys) + } } func TestCodexAccountIndexDoesNotMatchDifferentUsersInSameChatGPTAccount(t *testing.T) { @@ -333,35 +339,37 @@ func TestCodexAccountIndexDoesNotMatchDifferentUsersInSameChatGPTAccount(t *test "chatgpt_account_id": "team-1", "chatgpt_user_id": "user-1", "access_token": "token-1", + "refresh_token": "refresh-1", }, } index := buildCodexAccountIndex([]service.Account{existing}) - keys := buildCodexIdentityKeys("team-1", "user-2", "", "token-2") + keys := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-2", "refresh-2") if got, _ := index.Find(keys, "user-2"); got != nil { t.Fatalf("Find matched account ID %d for a different chatgpt_user_id in the same team", got.ID) } - keys = buildCodexIdentityKeys("team-1", "user-1", "", "token-2") + keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-2", "refresh-2") got, _ := index.Find(keys, "user-1") if got == nil || got.ID != existing.ID { t.Fatalf("Find by same chatgpt_user_id = %v, want account ID %d", got, existing.ID) } } -func TestCodexAccountIndexFallsBackToAccountKeyWhenUserIDMissing(t *testing.T) { - // 存量账号缺少 chatgpt_user_id:携带 user id 的重新导入应命中并更新(回填), - // 而不是创建重复账号。 +func TestCodexAccountIndexFallsBackToAccountKeyWhenRefreshTokenExistsAndUserIDMissing(t *testing.T) { + // 含 refresh_token 的常规导入沿用 a5638a4e 的兼容逻辑:存量账号缺少 + // chatgpt_user_id 时,携带 user id 的重新导入仍可命中并回填。 legacy := service.Account{ ID: 20, Credentials: map[string]any{ "chatgpt_account_id": "team-1", "access_token": "token-old", + "refresh_token": "refresh-old", }, } index := buildCodexAccountIndex([]service.Account{legacy}) - keys := buildCodexIdentityKeys("team-1", "user-1", "", "token-new") + keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "refresh-new") got, matchedKey := index.Find(keys, "user-1") if got == nil || got.ID != legacy.ID { t.Fatalf("Find legacy account without stored user id = %v, want account ID %d", got, legacy.ID) @@ -370,30 +378,59 @@ func TestCodexAccountIndexFallsBackToAccountKeyWhenUserIDMissing(t *testing.T) { t.Fatalf("matched key = %q, want account:team-1", matchedKey) } - // 反向:导入条目无法解析出 user id 时,仍应通过 account 键命中已有账号。 + // 反向:含 refresh_token 的导入条目无法解析出 user id 时,仍应通过 + // account 键命中已有账号,保持常规导入去重行为。 full := service.Account{ ID: 21, Credentials: map[string]any{ "chatgpt_account_id": "team-2", "chatgpt_user_id": "user-9", "access_token": "token-old", + "refresh_token": "refresh-old", }, } index = buildCodexAccountIndex([]service.Account{full}) - keys = buildCodexIdentityKeys("team-2", "", "", "token-opaque") + keys = buildCodexImportIdentityKeys("team-2", "", "", "token-opaque", "refresh-new") got, _ = index.Find(keys, "") if got == nil || got.ID != full.ID { t.Fatalf("Find by account key without entry user id = %v, want account ID %d", got, full.ID) } } +func TestCodexAccountIndexAccessTokenOnlyUsesTokenFingerprint(t *testing.T) { + existing := service.Account{ + ID: 22, + Credentials: map[string]any{ + "chatgpt_account_id": "team-1", + "chatgpt_user_id": "user-1", + "access_token": "token-old", + }, + } + index := buildCodexAccountIndex([]service.Account{existing}) + + keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "") + if got, matchedKey := index.Find(keys, "user-1"); got != nil { + t.Fatalf("accessToken-only import matched by %q despite different token: account ID %d", matchedKey, got.ID) + } + + keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-old", "") + got, matchedKey := index.Find(keys, "user-1") + if got == nil || got.ID != existing.ID { + t.Fatalf("Find accessToken-only duplicate by fingerprint = %v, want account ID %d", got, existing.ID) + } + if !strings.HasPrefix(matchedKey, "access:") { + t.Fatalf("matched key = %q, want access fingerprint", matchedKey) + } +} + func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) { legacy := service.Account{ ID: 30, Credentials: map[string]any{ "chatgpt_account_id": "team-1", "access_token": "token-legacy", + "refresh_token": "refresh-legacy", }, } member := service.Account{ @@ -402,10 +439,11 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) { "chatgpt_account_id": "team-1", "chatgpt_user_id": "user-2", "access_token": "token-member", + "refresh_token": "refresh-member", }, } - // 无论索引构建顺序如何,携带新 user id 的条目都应跳过 user-2 的账号、 + // 无论索引构建顺序如何,携带新 user id 的条目都应跳过 user-2 的账号, // 命中缺少 user id 的存量账号,而不是因单一候选被遮蔽而落空。 for _, accounts := range [][]service.Account{ {member, legacy}, @@ -413,7 +451,7 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) { } { index := buildCodexAccountIndex(accounts) - keys := buildCodexIdentityKeys("team-1", "user-1", "", "token-new") + keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "refresh-new") got, matchedKey := index.Find(keys, "user-1") if got == nil || got.ID != legacy.ID { t.Fatalf("Find with shared account key = %v, want legacy account ID %d", got, legacy.ID) @@ -422,7 +460,7 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) { t.Fatalf("matched key = %q, want account:team-1", matchedKey) } - keys = buildCodexIdentityKeys("team-1", "user-2", "", "token-new") + keys = buildCodexImportIdentityKeys("team-1", "user-2", "", "token-new", "refresh-new") got, matchedKey = index.Find(keys, "user-2") if got == nil || got.ID != member.ID { t.Fatalf("Find by user key = %v, want member account ID %d", got, member.ID) @@ -449,18 +487,19 @@ func TestCodexAccountIndexUpsertReplacesSameAccount(t *testing.T) { "chatgpt_account_id": "team-1", "chatgpt_user_id": "user-1", "access_token": "token-new", + "refresh_token": "refresh-new", }, } index.Add(backfilled) // 回填后同一账号在 account 键下应被原位替换而非残留旧副本: // 其他成员的条目不应再通过旧副本(无 user id)命中该账号。 - keys := buildCodexIdentityKeys("team-1", "user-2", "", "token-other") - if got, _ := index.Find(keys, "user-2"); got != nil { - t.Fatalf("stale candidate matched after upsert: account ID %d", got.ID) + keys := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-other", "refresh-other") + if got, matchedKey := index.Find(keys, "user-2"); got != nil { + t.Fatalf("stale candidate matched after upsert by %q: account ID %d", matchedKey, got.ID) } - keys = buildCodexIdentityKeys("team-1", "user-1", "", "token-other") + keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-other", "refresh-other") got, _ := index.Find(keys, "user-1") if got == nil || got.ID != backfilled.ID { t.Fatalf("Find after upsert = %v, want account ID %d", got, backfilled.ID) @@ -472,28 +511,421 @@ func TestCodexAccountIndexUpsertReplacesSameAccount(t *testing.T) { func TestCodexIdentitySeenDistinguishesTeamMembers(t *testing.T) { seen := map[string]codexSeenIdentity{} - member1 := buildCodexIdentityKeys("team-1", "user-1", "", "token-1") + member1 := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-1", "refresh-1") markCodexIdentitySeen(seen, member1, 1, "user-1") - member2 := buildCodexIdentityKeys("team-1", "user-2", "", "token-2") + member2 := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-2", "refresh-2") if index, ok := firstSeenCodexIdentity(seen, member2, "user-2"); ok { t.Fatalf("different team member treated as duplicate of entry %d", index) } - again := buildCodexIdentityKeys("team-1", "user-1", "", "token-3") + again := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-3", "refresh-3") index, ok := firstSeenCodexIdentity(seen, again, "user-1") if !ok || index != 1 { t.Fatalf("same user re-entry dedup = (%d, %v), want (1, true)", index, ok) } - // 无 user id 的条目与已见同 account 条目视为重复(保守跳过,与既有行为一致)。 - opaque := buildCodexIdentityKeys("team-1", "", "", "token-4") + // 无 user id 的条目不应因共享 account id 与已见团队成员互相去重; + // 只有相同 access token 指纹才视为重复。 + opaque := buildCodexImportIdentityKeys("team-1", "", "", "token-4", "") index, ok = firstSeenCodexIdentity(seen, opaque, "") - if !ok || index != 1 { - t.Fatalf("entry without user id dedup = (%d, %v), want (1, true)", index, ok) + if ok { + t.Fatalf("entry without user id dedup = (%d, %v), want no match", index, ok) } } +func TestNormalizeCodexImportUsesJWTSubForAccessTokenOnlyIdentity(t *testing.T) { + accessToken := buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{ + "sub": "user-from-access-token", + "https://api.openai.com/auth": map[string]any{ + "chatgpt_account_id": "workspace-1", + }, + }) + + item, err := normalizeCodexImportEntry(codexImportEntry{Index: 1, Value: accessToken}) + if err != nil { + t.Fatalf("normalizeCodexImportEntry error = %v", err) + } + if item.UserID != "user-from-access-token" { + t.Fatalf("UserID = %q, want JWT sub", item.UserID) + } + if len(item.IdentityKeys) != 1 || !strings.HasPrefix(item.IdentityKeys[0], "access:") { + t.Fatalf("IdentityKeys = %v, want access fingerprint only for accessToken-only import", item.IdentityKeys) + } + if got := item.Credentials["chatgpt_user_id"]; got != "user-from-access-token" { + t.Fatalf("credential chatgpt_user_id = %v, want JWT sub", got) + } +} + +func TestImportCodexSessionsAccessTokenOnlySameWorkspaceDifferentUsersCreatesTwoAccounts(t *testing.T) { + svc := newCodexImportMemoryAdminService(nil) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: buildCodexAccessOnlyImportValue(t, "workspace-1", "user-1")}, + {Index: 2, Value: buildCodexAccessOnlyImportValue(t, "workspace-1", "user-2")}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 2 || result.Updated != 0 || result.Skipped != 0 || result.Failed != 0 { + t.Fatalf("result = %+v, want two created accounts", result) + } + if len(svc.createdAccounts) != 2 { + t.Fatalf("created accounts = %d, want 2", len(svc.createdAccounts)) + } + if svc.createdAccounts[0].Credentials["chatgpt_user_id"] == svc.createdAccounts[1].Credentials["chatgpt_user_id"] { + t.Fatalf("created accounts share user id: %v", svc.createdAccounts) + } +} + +func TestImportCodexSessionsAccessTokenOnlySameWorkspaceAndUserDifferentTokensCreatesTwoAccounts(t *testing.T) { + svc := newCodexImportMemoryAdminService(nil) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: map[string]any{ + "access_token": buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{ + "sub": "shared-user", + "jti": "token-1", + "https://api.openai.com/auth": map[string]any{ + "chatgpt_account_id": "workspace-1", + }, + }), + }}, + {Index: 2, Value: map[string]any{ + "access_token": buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{ + "sub": "shared-user", + "jti": "token-2", + "https://api.openai.com/auth": map[string]any{ + "chatgpt_account_id": "workspace-1", + }, + }), + }}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 2 || result.Updated != 0 || result.Skipped != 0 || result.Failed != 0 { + t.Fatalf("result = %+v, want two created accounts", result) + } + if len(svc.createdAccounts) != 2 { + t.Fatalf("created accounts = %d, want 2", len(svc.createdAccounts)) + } +} + +func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) { + existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour)) + svc := newCodexImportMemoryAdminService([]service.Account{{ + ID: 10, + Name: "existing", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "workspace-1", + "chatgpt_user_id": "user-1", + "access_token": existingToken, + }, + }}) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: map[string]any{"access_token": existingToken}}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 0 || result.Updated != 1 || result.Failed != 0 { + t.Fatalf("result = %+v, want one updated account", result) + } + if len(svc.createdAccounts) != 0 { + t.Fatalf("created accounts = %d, want 0", len(svc.createdAccounts)) + } + if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 { + t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts) + } +} + +func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) { + oldToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "old-token", time.Now().Add(time.Hour)) + newToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "new-token", time.Now().Add(time.Hour)) + svc := newCodexImportMemoryAdminService([]service.Account{{ + ID: 12, + Name: "existing", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "workspace-1", + "chatgpt_user_id": "user-1", + "access_token": oldToken, + }, + }}) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: map[string]any{ + "access_token": newToken, + "refresh_token": "refresh-new", + }}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 0 || result.Updated != 1 || result.Failed != 0 { + t.Fatalf("result = %+v, want one updated account", result) + } + if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 12 { + t.Fatalf("updated accounts = %+v, want account 12", svc.updatedAccounts) + } + if got := svc.updatedAccounts[0].input.Credentials["refresh_token"]; got != "refresh-new" { + t.Fatalf("updated refresh_token = %v, want refresh-new", got) + } +} + +func TestImportCodexSessionsAccessTokenOnlyPreservesExistingRefreshToken(t *testing.T) { + existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour)) + svc := newCodexImportMemoryAdminService([]service.Account{{ + ID: 13, + Name: "existing", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "workspace-1", + "chatgpt_user_id": "user-1", + "access_token": existingToken, + "refresh_token": "refresh-old", + "client_id": "client-old", + }, + }}) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: map[string]any{"access_token": existingToken}}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 0 || result.Updated != 1 || result.Failed != 0 { + t.Fatalf("result = %+v, want one updated account", result) + } + update := svc.updatedAccounts[0].input + if got := update.Credentials["refresh_token"]; got != "refresh-old" { + t.Fatalf("refresh_token = %v, want refresh-old", got) + } + if got := update.Credentials["client_id"]; got != "client-old" { + t.Fatalf("client_id = %v, want client-old", got) + } + if update.ExpiresAt != nil { + t.Fatalf("ExpiresAt = %v, want nil to preserve OAuth account expiry", *update.ExpiresAt) + } + if update.AutoPauseOnExpired != nil { + t.Fatalf("AutoPauseOnExpired = %v, want nil to preserve OAuth account scheduling", *update.AutoPauseOnExpired) + } +} + +func TestImportCodexSessionsBatchOldAccessTokenDoesNotRollbackRefreshToken(t *testing.T) { + oldToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "old-token", time.Now().Add(time.Hour)) + newToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "new-token", time.Now().Add(time.Hour)) + svc := newCodexImportMemoryAdminService([]service.Account{{ + ID: 14, + Name: "existing", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "workspace-1", + "chatgpt_user_id": "user-1", + "access_token": oldToken, + "refresh_token": "refresh-old", + }, + }}) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: map[string]any{ + "access_token": newToken, + "refresh_token": "refresh-new", + }}, + {Index: 2, Value: map[string]any{"access_token": oldToken}}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Updated != 1 || result.Created != 1 || result.Failed != 0 { + t.Fatalf("result = %+v, want first item updated and stale access token created separately", result) + } + if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 14 { + t.Fatalf("updated accounts = %+v, want account 14 updated once", svc.updatedAccounts) + } + stored, err := svc.GetAccount(context.Background(), 14) + if err != nil { + t.Fatalf("GetAccount error = %v", err) + } + if got := stored.Credentials["access_token"]; got != newToken { + t.Fatalf("stored access_token rolled back = %v, want new token", got) + } + if got := stored.Credentials["refresh_token"]; got != "refresh-new" { + t.Fatalf("stored refresh_token = %v, want refresh-new", got) + } +} + +func TestImportCodexSessionsWithRefreshTokenKeepsExistingDedup(t *testing.T) { + existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour)) + svc := newCodexImportMemoryAdminService([]service.Account{{ + ID: 11, + Name: "existing", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "workspace-1", + "chatgpt_user_id": "user-1", + "access_token": existingToken, + "refresh_token": "refresh-old", + }, + }}) + handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} + entries := []codexImportEntry{ + {Index: 1, Value: buildCodexRefreshImportValue(t, "workspace-1", "user-1", "refresh-new")}, + } + + result, err := handler.importCodexSessions(context.Background(), req, entries) + if err != nil { + t.Fatalf("importCodexSessions error = %v", err) + } + if result.Created != 0 || result.Updated != 1 || result.Failed != 0 { + t.Fatalf("result = %+v, want one updated account", result) + } + if got := svc.updatedAccounts[0].input.Credentials["refresh_token"]; got != "refresh-new" { + t.Fatalf("updated refresh_token = %v, want refresh-new", got) + } +} + +type codexImportMemoryAdminService struct { + *stubAdminService + nextID int64 + updatedAccounts []struct { + id int64 + input *service.UpdateAccountInput + } +} + +func newCodexImportMemoryAdminService(accounts []service.Account) *codexImportMemoryAdminService { + stub := newStubAdminService() + stub.accounts = append([]service.Account(nil), accounts...) + return &codexImportMemoryAdminService{ + stubAdminService: stub, + nextID: 100, + } +} + +func (s *codexImportMemoryAdminService) CreateAccount(ctx context.Context, input *service.CreateAccountInput) (*service.Account, error) { + s.createdAccounts = append(s.createdAccounts, input) + if s.createAccountErr != nil { + return nil, s.createAccountErr + } + account := service.Account{ + ID: s.nextID, + Name: input.Name, + Platform: input.Platform, + Type: input.Type, + Status: service.StatusActive, + Credentials: cloneCodexImportTestMap(input.Credentials), + Extra: cloneCodexImportTestMap(input.Extra), + } + s.nextID++ + s.accounts = append(s.accounts, account) + return &account, nil +} + +func (s *codexImportMemoryAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) { + s.updatedAccounts = append(s.updatedAccounts, struct { + id int64 + input *service.UpdateAccountInput + }{id: id, input: input}) + if s.updateAccountErr != nil { + return nil, s.updateAccountErr + } + for idx := range s.accounts { + if s.accounts[idx].ID == id { + s.accounts[idx].Credentials = cloneCodexImportTestMap(input.Credentials) + s.accounts[idx].Extra = cloneCodexImportTestMap(input.Extra) + return &s.accounts[idx], nil + } + } + account := service.Account{ID: id, Status: service.StatusActive, Credentials: cloneCodexImportTestMap(input.Credentials)} + return &account, nil +} + +func (s *codexImportMemoryAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) { + for idx := range s.accounts { + if s.accounts[idx].ID == id { + return &s.accounts[idx], nil + } + } + return s.stubAdminService.GetAccount(ctx, id) +} + +func buildCodexAccessOnlyImportValue(t *testing.T, accountID, userID string) map[string]any { + t.Helper() + return map[string]any{ + "access_token": buildCodexAccessToken(t, accountID, userID, time.Now().Add(time.Hour)), + } +} + +func buildCodexRefreshImportValue(t *testing.T, accountID, userID, refreshToken string) map[string]any { + t.Helper() + return map[string]any{ + "access_token": buildCodexAccessToken(t, accountID, userID, time.Now().Add(time.Hour)), + "refresh_token": refreshToken, + } +} + +func buildCodexAccessToken(t *testing.T, accountID, userID string, exp time.Time) string { + t.Helper() + return buildCodexAccessTokenWithJTI(t, accountID, userID, "", exp) +} + +func buildCodexAccessTokenWithJTI(t *testing.T, accountID, userID, jti string, exp time.Time) string { + t.Helper() + claims := map[string]any{ + "sub": userID, + "https://api.openai.com/auth": map[string]any{ + "chatgpt_account_id": accountID, + }, + } + if jti != "" { + claims["jti"] = jti + } + return buildCodexImportTestJWT(t, exp, claims) +} + +func cloneCodexImportTestMap(input map[string]any) map[string]any { + if input == nil { + return nil + } + out := make(map[string]any, len(input)) + for key, value := range input { + out[key] = value + } + return out +} + +func boolPtr(v bool) *bool { + return &v +} + func buildCodexImportTestJWT(t *testing.T, exp time.Time, extraClaims map[string]any) string { t.Helper() header := map[string]any{ diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 8d4a5dfea6..8c91245fbf 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -171,13 +171,29 @@ type CheckMixedChannelRequest struct { // AccountWithConcurrency extends Account with real-time concurrency info type AccountWithConcurrency struct { *dto.Account - CurrentConcurrency int `json:"current_concurrency"` + CurrentConcurrency int `json:"current_concurrency"` + SchedulerScore *AccountSchedulerScore `json:"scheduler_score,omitempty"` + SchedulerScores []AccountSchedulerGroupScore `json:"scheduler_scores,omitempty"` // 以下字段仅对 Anthropic OAuth/SetupToken 账号有效,且仅在启用相应功能时返回 CurrentWindowCost *float64 `json:"current_window_cost,omitempty"` // 当前窗口费用 ActiveSessions *int `json:"active_sessions,omitempty"` // 当前活跃会话数 CurrentRPM *int `json:"current_rpm,omitempty"` // 当前分钟 RPM 计数 } +type AccountSchedulerScore struct { + BaseScore float64 `json:"base_score"` + StickyScore float64 `json:"sticky_score"` + StickyScoreInfinity bool `json:"sticky_score_infinity"` + StickyWeightedEnabled bool `json:"sticky_weighted_enabled"` +} + +type AccountSchedulerGroupScore struct { + GroupID *int64 `json:"group_id"` + GroupName string `json:"group_name,omitempty"` + GroupPriority *int `json:"group_priority,omitempty"` + AccountSchedulerScore +} + const accountListGroupUngroupedQueryValue = "ungrouped" func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, account *service.Account) AccountWithConcurrency { @@ -226,6 +242,232 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac return item } +// scoreOpenAIAccountSchedulerPool 对池内 OpenAI 账号计算调度分数快照。 +// loadMap 为共享的账号负载数据(含池内全部账号即可,多余条目无害);传 nil 时自行批查。 +func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account, loadMap map[int64]*service.AccountLoadInfo) map[int64]AccountSchedulerScore { + if len(accounts) == 0 { + return nil + } + + openAIAccounts := make([]*service.Account, 0, len(accounts)) + for i := range accounts { + account := &accounts[i] + if account.Platform != service.PlatformOpenAI { + continue + } + openAIAccounts = append(openAIAccounts, account) + } + if len(openAIAccounts) == 0 { + return nil + } + + if loadMap == nil { + loadMap = h.fetchOpenAIAccountLoadMap(ctx, openAIAccounts) + } + + var scores map[int64]service.OpenAIAccountSchedulerScoreSnapshot + if h.rateLimitService != nil { + scores = h.rateLimitService.BuildOpenAIAccountSchedulerScoreSnapshot(ctx, openAIAccounts, loadMap) + } else { + scores = service.BuildOpenAIAccountSchedulerScoreSnapshot(openAIAccounts, loadMap) + } + result := make(map[int64]AccountSchedulerScore, len(scores)) + for accountID, score := range scores { + result[accountID] = AccountSchedulerScore{ + BaseScore: score.BaseScore, + StickyScore: score.StickyScore, + StickyScoreInfinity: score.StickyScoreInfinity, + StickyWeightedEnabled: score.StickyWeightedEnabled, + } + } + return result +} + +// fetchOpenAIAccountLoadMap 一次性批查给定 OpenAI 账号的负载数据; +// 失败时记录日志并返回空表(分数按零负载计算,属可接受降级)。 +func (h *AccountHandler) fetchOpenAIAccountLoadMap(ctx context.Context, openAIAccounts []*service.Account) map[int64]*service.AccountLoadInfo { + loadMap := map[int64]*service.AccountLoadInfo{} + if h.concurrencyService == nil || len(openAIAccounts) == 0 { + return loadMap + } + seen := make(map[int64]struct{}, len(openAIAccounts)) + loadReq := make([]service.AccountWithConcurrency, 0, len(openAIAccounts)) + for _, account := range openAIAccounts { + if account == nil { + continue + } + if _, ok := seen[account.ID]; ok { + continue + } + seen[account.ID] = struct{}{} + loadReq = append(loadReq, service.AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err != nil { + slog.Warn("openai_scheduler_score_load_batch_failed", "error", err) + } else if batchLoad != nil { + loadMap = batchLoad + } + return loadMap +} + +func (h *AccountHandler) buildOpenAIAccountSchedulerScores( + ctx context.Context, + accounts []service.Account, + filterPool []service.Account, +) (map[int64]*AccountSchedulerScore, map[int64][]AccountSchedulerGroupScore) { + if len(accounts) == 0 { + return nil, nil + } + if len(filterPool) == 0 { + filterPool = accounts + } + + pageOpenAIAccountIDs := make(map[int64]struct{}) + groupIDs := make(map[int64]struct{}) + for i := range accounts { + account := &accounts[i] + if account.Platform != service.PlatformOpenAI { + continue + } + pageOpenAIAccountIDs[account.ID] = struct{}{} + if len(account.AccountGroups) == 0 && len(account.GroupIDs) == 0 { + continue + } + for _, accountGroup := range account.AccountGroups { + if accountGroup.GroupID > 0 { + groupIDs[accountGroup.GroupID] = struct{}{} + } + } + for _, groupID := range account.GroupIDs { + if groupID > 0 { + groupIDs[groupID] = struct{}{} + } + } + } + if len(pageOpenAIAccountIDs) == 0 { + return nil, nil + } + + // 先取各分组池,再对"过滤池 ∪ 分组池"的账号并集做一次负载批查, + // 避免每个池各查一次 Redis 的 N+1。 + groupIDList := make([]int64, 0, len(groupIDs)) + for groupID := range groupIDs { + groupIDList = append(groupIDList, groupID) + } + sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] }) + + groupPools := make(map[int64][]service.Account, len(groupIDList)) + if h.adminService != nil { + for _, groupID := range groupIDList { + gid := groupID + pool, err := h.adminService.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid) + if err != nil { + slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err) + continue + } + groupPools[gid] = pool + } + } + + loadUnion := make([]*service.Account, 0, len(filterPool)) + collectOpenAIAccounts := func(pool []service.Account) { + for i := range pool { + if pool[i].Platform == service.PlatformOpenAI { + loadUnion = append(loadUnion, &pool[i]) + } + } + } + collectOpenAIAccounts(filterPool) + for _, pool := range groupPools { + collectOpenAIAccounts(pool) + } + loadMap := h.fetchOpenAIAccountLoadMap(ctx, loadUnion) + + baseScores := make(map[int64]*AccountSchedulerScore) + for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool, loadMap) { + copiedScore := score + baseScores[accountID] = &copiedScore + } + + groupScoresByAccount := make(map[int64][]AccountSchedulerGroupScore) + scoreGroupPool := func(groupID *int64, groupNameByID map[int64]string, groupPriorityByAccount map[int64]int, pool []service.Account) { + if len(pool) == 0 { + return + } + scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool, loadMap) + for accountID, schedulerScore := range scores { + if _, ok := pageOpenAIAccountIDs[accountID]; !ok { + continue + } + groupScore := AccountSchedulerGroupScore{ + GroupID: groupID, + AccountSchedulerScore: schedulerScore, + } + if groupID != nil { + groupScore.GroupName = groupNameByID[*groupID] + if priority, ok := groupPriorityByAccount[accountID]; ok { + groupScore.GroupPriority = &priority + } + } + groupScoresByAccount[accountID] = append(groupScoresByAccount[accountID], groupScore) + } + } + + for _, groupID := range groupIDList { + gid := groupID + pool, ok := groupPools[gid] + if !ok { + continue + } + groupNameByID := make(map[int64]string) + groupPriorityByAccount := make(map[int64]int) + for i := range pool { + account := &pool[i] + for _, accountGroup := range account.AccountGroups { + if accountGroup.GroupID != gid { + continue + } + groupPriorityByAccount[account.ID] = accountGroup.Priority + if accountGroup.Group != nil { + groupNameByID[gid] = accountGroup.Group.Name + } + } + } + scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool) + } + + for accountID := range groupScoresByAccount { + sort.SliceStable(groupScoresByAccount[accountID], func(i, j int) bool { + left := groupScoresByAccount[accountID][i] + right := groupScoresByAccount[accountID][j] + return *left.GroupID < *right.GroupID + }) + } + return baseScores, groupScoresByAccount +} + +func (h *AccountHandler) listAccountSchedulerScoreFilterPool( + ctx context.Context, + platform, accountType, status, search string, + groupID int64, + privacyMode string, +) []service.Account { + if h.adminService == nil || (platform != "" && platform != service.PlatformOpenAI) { + return nil + } + // 池只用于 OpenAI 分数计算(非 OpenAI 账号会在打分时被丢弃), + // 无论列表页平台过滤为何,查询一律限定 openai,避免无过滤时全表扫描。 + accounts, err := h.adminService.ListAccountsForSchedulerScoreFilter(ctx, service.PlatformOpenAI, accountType, status, search, groupID, privacyMode) + if err != nil { + slog.Warn("openai_scheduler_filter_score_pool_failed", "error", err) + return nil + } + return accounts +} + // List handles listing all accounts with pagination // GET /api/v1/admin/accounts func (h *AccountHandler) List(c *gin.Context) { @@ -278,6 +520,20 @@ func (h *AccountHandler) List(c *gin.Context) { var windowCosts map[int64]float64 var activeSessions map[int64]int var rpmCounts map[int64]int + // 仅当前页存在 OpenAI 账号时才计算调度分数,避免为空结果付出池查询开销。 + var schedulerScores map[int64]*AccountSchedulerScore + var schedulerGroupScores map[int64][]AccountSchedulerGroupScore + pageHasOpenAIAccounts := false + for i := range accounts { + if accounts[i].Platform == service.PlatformOpenAI { + pageHasOpenAIAccounts = true + break + } + } + if pageHasOpenAIAccounts { + schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode) + schedulerScores, schedulerGroupScores = h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool) + } // 始终获取并发数(Redis ZCARD,极低开销) if h.concurrencyService != nil { @@ -358,6 +614,8 @@ func (h *AccountHandler) List(c *gin.Context) { item := AccountWithConcurrency{ Account: dto.AccountFromService(acc), CurrentConcurrency: concurrencyCounts[acc.ID], + SchedulerScore: schedulerScores[acc.ID], + SchedulerScores: schedulerGroupScores[acc.ID], } // 添加窗口费用(仅当启用时) diff --git a/backend/internal/handler/admin/account_handler_list_test.go b/backend/internal/handler/admin/account_handler_list_test.go index 4d628365df..29e36ad865 100644 --- a/backend/internal/handler/admin/account_handler_list_test.go +++ b/backend/internal/handler/admin/account_handler_list_test.go @@ -8,6 +8,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -50,3 +51,222 @@ func TestAccountHandlerListIncludesCreatedAt(t *testing.T) { _, offset := parsed.Zone() require.Equal(t, 0, offset) } + +func TestAccountHandlerListReturnsSchedulerScoresPerGroup(t *testing.T) { + router, adminSvc := setupAccountListRouter() + now := time.Now().UTC() + groupID := int64(41) + adminSvc.accounts = []service.Account{ + { + ID: 101, + Name: "account-high-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 1, + AccountGroups: []service.AccountGroup{ + {AccountID: 101, GroupID: groupID, Priority: 100, Group: &service.Group{ID: groupID, Name: "openai"}}, + }, + GroupIDs: []int64{groupID}, + CreatedAt: now, + UpdatedAt: now, + }, + { + ID: 102, + Name: "account-low-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 100000, + AccountGroups: []service.AccountGroup{ + {AccountID: 102, GroupID: groupID, Priority: 1, Group: &service.Group{ID: groupID, Name: "openai"}}, + }, + GroupIDs: []int64{groupID}, + CreatedAt: now, + UpdatedAt: now, + }, + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=20&platform=openai", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var payload struct { + Data struct { + Items []struct { + ID int64 `json:"id"` + SchedulerScore struct { + BaseScore float64 `json:"base_score"` + } `json:"scheduler_score"` + SchedulerScores []struct { + GroupID *int64 `json:"group_id"` + GroupName string `json:"group_name"` + GroupPriority *int `json:"group_priority"` + BaseScore float64 `json:"base_score"` + } `json:"scheduler_scores"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload)) + require.Len(t, payload.Data.Items, 2) + + var high, low *struct { + ID int64 `json:"id"` + SchedulerScore struct { + BaseScore float64 `json:"base_score"` + } `json:"scheduler_score"` + SchedulerScores []struct { + GroupID *int64 `json:"group_id"` + GroupName string `json:"group_name"` + GroupPriority *int `json:"group_priority"` + BaseScore float64 `json:"base_score"` + } `json:"scheduler_scores"` + } + for i := range payload.Data.Items { + item := &payload.Data.Items[i] + switch item.ID { + case 101: + high = item + case 102: + low = item + } + } + require.NotNil(t, high) + require.NotNil(t, low) + require.Len(t, high.SchedulerScores, 1) + require.Len(t, low.SchedulerScores, 1) + require.Equal(t, groupID, *high.SchedulerScores[0].GroupID) + require.Equal(t, "openai", high.SchedulerScores[0].GroupName) + require.Equal(t, 100, *high.SchedulerScores[0].GroupPriority) + require.Equal(t, 1, *low.SchedulerScores[0].GroupPriority) + require.Greater(t, high.SchedulerScores[0].BaseScore, low.SchedulerScores[0].BaseScore) +} + +func TestAccountHandlerListKeepsSchedulerScoreScopedToFilter(t *testing.T) { + router, adminSvc := setupAccountListRouter() + now := time.Now().UTC() + groupID := int64(42) + visibleAccount := service.Account{ + ID: 201, + Name: "visible-low-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 100000, + AccountGroups: []service.AccountGroup{ + {AccountID: 201, GroupID: groupID, Priority: 1, Group: &service.Group{ID: groupID, Name: "openai"}}, + }, + GroupIDs: []int64{groupID}, + CreatedAt: now, + UpdatedAt: now, + } + hiddenGroupPeer := service.Account{ + ID: 202, + Name: "hidden-high-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 1, + AccountGroups: []service.AccountGroup{ + {AccountID: 202, GroupID: groupID, Priority: 2, Group: &service.Group{ID: groupID, Name: "openai"}}, + }, + GroupIDs: []int64{groupID}, + CreatedAt: now, + UpdatedAt: now, + } + adminSvc.accounts = []service.Account{visibleAccount} + adminSvc.accountSchedulerScoreFilterAccounts = []service.Account{visibleAccount, hiddenGroupPeer} + adminSvc.openAISchedulerScorePoolAccounts = []service.Account{visibleAccount, hiddenGroupPeer} + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=1&platform=openai", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var payload struct { + Data struct { + Items []struct { + ID int64 `json:"id"` + SchedulerScore struct { + BaseScore float64 `json:"base_score"` + } `json:"scheduler_score"` + SchedulerScores []struct { + GroupID *int64 `json:"group_id"` + BaseScore float64 `json:"base_score"` + } `json:"scheduler_scores"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload)) + require.Len(t, payload.Data.Items, 1) + item := payload.Data.Items[0] + require.Equal(t, int64(201), item.ID) + require.Len(t, item.SchedulerScores, 1) + require.Equal(t, groupID, *item.SchedulerScores[0].GroupID) + require.Equal(t, item.SchedulerScores[0].BaseScore, item.SchedulerScore.BaseScore) +} + +func TestAccountHandlerListSchedulerScoreIgnoresPagination(t *testing.T) { + router, adminSvc := setupAccountListRouter() + now := time.Now().UTC() + visibleAccount := service.Account{ + ID: 301, + Name: "visible-low-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 100000, + CreatedAt: now, + UpdatedAt: now, + } + hiddenFilterPeer := service.Account{ + ID: 302, + Name: "hidden-high-priority", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 10, + Priority: 1, + CreatedAt: now, + UpdatedAt: now, + } + adminSvc.accounts = []service.Account{visibleAccount} + adminSvc.accountSchedulerScoreFilterAccounts = []service.Account{visibleAccount, hiddenFilterPeer} + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=1&platform=openai", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var payload struct { + Data struct { + Items []struct { + ID int64 `json:"id"` + SchedulerScore struct { + BaseScore float64 `json:"base_score"` + } `json:"scheduler_score"` + SchedulerScores []struct { + GroupID *int64 `json:"group_id"` + BaseScore float64 `json:"base_score"` + } `json:"scheduler_scores"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload)) + require.Len(t, payload.Data.Items, 1) + require.Equal(t, int64(301), payload.Data.Items[0].ID) + require.Less(t, payload.Data.Items[0].SchedulerScore.BaseScore, 3.75) + require.Empty(t, payload.Data.Items[0].SchedulerScores) +} diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go index f7bab13253..187925e33a 100644 --- a/backend/internal/handler/admin/admin_service_stub_test.go +++ b/backend/internal/handler/admin/admin_service_stub_test.go @@ -10,27 +10,29 @@ import ( ) type stubAdminService struct { - users []service.User - apiKeys []service.APIKey - groups []service.Group - accounts []service.Account - proxies []service.Proxy - proxyCounts []service.ProxyWithAccountCount - redeems []service.RedeemCode - boundAuthIdentity *service.AdminBindAuthIdentityInput - boundAuthIdentityFor int64 - createdAccounts []*service.CreateAccountInput - createdProxies []*service.CreateProxyInput - updatedProxyIDs []int64 - updatedProxies []*service.UpdateProxyInput - testedProxyIDs []int64 - getUserErr error - createAccountErr error - createSparkShadowErr error - updateAccountErr error - bulkUpdateAccountErr error - checkMixedErr error - lastMixedCheck struct { + users []service.User + apiKeys []service.APIKey + groups []service.Group + accounts []service.Account + accountSchedulerScoreFilterAccounts []service.Account + openAISchedulerScorePoolAccounts []service.Account + proxies []service.Proxy + proxyCounts []service.ProxyWithAccountCount + redeems []service.RedeemCode + boundAuthIdentity *service.AdminBindAuthIdentityInput + boundAuthIdentityFor int64 + createdAccounts []*service.CreateAccountInput + createdProxies []*service.CreateProxyInput + updatedProxyIDs []int64 + updatedProxies []*service.UpdateProxyInput + testedProxyIDs []int64 + getUserErr error + createAccountErr error + createSparkShadowErr error + updateAccountErr error + bulkUpdateAccountErr error + checkMixedErr error + lastMixedCheck struct { accountID int64 platform string groupIDs []int64 @@ -329,7 +331,56 @@ func (s *stubAdminService) ListAccounts(ctx context.Context, page, pageSize int, s.lastListAccounts.sortBy = sortBy s.lastListAccounts.sortOrder = sortOrder s.lastListAccounts.calls++ - return s.accounts, int64(len(s.accounts)), nil + accounts := s.accounts + total := len(accounts) + if page < 1 { + page = 1 + } + if pageSize < 1 { + pageSize = total + } + start := (page - 1) * pageSize + if start >= total { + return []service.Account{}, int64(total), nil + } + end := start + pageSize + if end > total { + end = total + } + return accounts[start:end], int64(total), nil +} + +func (s *stubAdminService) ListAccountsForSchedulerScoreFilter(_ context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) { + if s.accountSchedulerScoreFilterAccounts != nil { + return s.accountSchedulerScoreFilterAccounts, nil + } + return s.accounts, nil +} + +func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ context.Context, groupID *int64) ([]service.Account, error) { + accounts := s.openAISchedulerScorePoolAccounts + if accounts == nil { + accounts = s.accounts + } + out := make([]service.Account, 0, len(accounts)) + for _, account := range accounts { + if account.Platform != service.PlatformOpenAI || !account.IsSchedulable() { + continue + } + if groupID == nil { + if len(account.AccountGroups) == 0 && len(account.GroupIDs) == 0 { + out = append(out, account) + } + continue + } + for _, accountGroup := range account.AccountGroups { + if accountGroup.GroupID == *groupID { + out = append(out, account) + break + } + } + } + return out, nil } func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) { diff --git a/backend/internal/handler/admin/ops_handler.go b/backend/internal/handler/admin/ops_handler.go index b9558b97b1..e820aef0c6 100644 --- a/backend/internal/handler/admin/ops_handler.go +++ b/backend/internal/handler/admin/ops_handler.go @@ -73,6 +73,13 @@ func NewOpsHandler(opsService *service.OpsService) *OpsHandler { } // GetErrorLogs lists ops error logs. +// applyOpsErrorSortParams reads sort_by/sort_order query params into the filter. +// Column whitelist and order normalization live in the repository; unknown +// values degrade to the default (created_at DESC), mirroring the usage list. +func applyOpsErrorSortParams(c *gin.Context, filter *service.OpsErrorLogFilter) { + filter.SetSort(c.Query("sort_by"), c.Query("sort_order")) +} + // GET /api/v1/admin/ops/errors func (h *OpsHandler) GetErrorLogs(c *gin.Context) { if h.opsService == nil { @@ -114,10 +121,17 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) { // buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。 filter.Model = strings.TrimSpace(c.Query("model")) - // Force request errors: client-visible status >= 400. - // buildOpsErrorLogsWhere already applies this for non-upstream phase. - if strings.EqualFold(strings.TrimSpace(filter.Phase), "upstream") { - filter.Phase = "" + // 请求错误语义:client-visible status>=400 守卫恒生效(未设 + // IncludeRecoveredUpstream 时 phase=upstream 不再绕过守卫),故 + // phase=upstream 作为普通过滤条件保留——此前这里清空该值,导致 + // 错误类型下拉选「上游」等于不过滤。 + + // 分类(用户侧粗分类码)→ phase/type ANY 条件,与用户端 /usage/errors 同一映射; + // 未知分类返回空切片 = 不过滤。与 phase 参数可同时设置(AND 语义)。 + if cat := strings.TrimSpace(c.Query("category")); cat != "" { + phases, types := service.CategoryToFilter(cat) + filter.ErrorPhasesAny = phases + filter.ErrorTypesAny = types } if platform := strings.TrimSpace(c.Query("platform")); platform != "" { @@ -187,6 +201,8 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) { filter.StatusCodes = out } + applyOpsErrorSortParams(c, filter) + result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter) if err != nil { response.ErrorFrom(c, err) @@ -234,10 +250,17 @@ func (h *OpsHandler) ListRequestErrors(c *gin.Context) { // buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。 filter.Model = strings.TrimSpace(c.Query("model")) - // Force request errors: client-visible status >= 400. - // buildOpsErrorLogsWhere already applies this for non-upstream phase. - if strings.EqualFold(strings.TrimSpace(filter.Phase), "upstream") { - filter.Phase = "" + // 请求错误语义:client-visible status>=400 守卫恒生效(未设 + // IncludeRecoveredUpstream 时 phase=upstream 不再绕过守卫),故 + // phase=upstream 作为普通过滤条件保留——此前这里清空该值,导致 + // 错误类型下拉选「上游」等于不过滤。 + + // 分类(用户侧粗分类码)→ phase/type ANY 条件,与用户端 /usage/errors 同一映射; + // 未知分类返回空切片 = 不过滤。与 phase 参数可同时设置(AND 语义)。 + if cat := strings.TrimSpace(c.Query("category")); cat != "" { + phases, types := service.CategoryToFilter(cat) + filter.ErrorPhasesAny = phases + filter.ErrorTypesAny = types } if platform := strings.TrimSpace(c.Query("platform")); platform != "" { @@ -291,6 +314,8 @@ func (h *OpsHandler) ListRequestErrors(c *gin.Context) { filter.StatusCodes = out } + applyOpsErrorSortParams(c, filter) + result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter) if err != nil { response.ErrorFrom(c, err) @@ -362,6 +387,8 @@ func (h *OpsHandler) ListRequestErrorUpstreamErrors(c *gin.Context) { } filter.View = "all" filter.Phase = "upstream" + // 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。 + filter.IncludeRecoveredUpstream = true filter.Owner = "provider" filter.Source = strings.TrimSpace(c.Query("error_source")) filter.Query = strings.TrimSpace(c.Query("q")) @@ -377,6 +404,8 @@ func (h *OpsHandler) ListRequestErrorUpstreamErrors(c *gin.Context) { filter.ClientRequestID = clientRequestID } + applyOpsErrorSortParams(c, filter) + result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter) if err != nil { response.ErrorFrom(c, err) @@ -442,6 +471,8 @@ func (h *OpsHandler) ListUpstreamErrors(c *gin.Context) { filter.View = parseOpsViewParam(c) filter.Phase = "upstream" + // 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。 + filter.IncludeRecoveredUpstream = true filter.Owner = "provider" filter.Source = strings.TrimSpace(c.Query("error_source")) filter.Query = strings.TrimSpace(c.Query("q")) @@ -497,6 +528,8 @@ func (h *OpsHandler) ListUpstreamErrors(c *gin.Context) { filter.StatusCodes = out } + applyOpsErrorSortParams(c, filter) + result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter) if err != nil { response.ErrorFrom(c, err) diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 0f50ab723b..529e46c575 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -119,188 +119,211 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { } payload := dto.SystemSettings{ - RegistrationEnabled: settings.RegistrationEnabled, - EmailVerifyEnabled: settings.EmailVerifyEnabled, - RegistrationEmailSuffixWhitelist: settings.RegistrationEmailSuffixWhitelist, - PromoCodeEnabled: settings.PromoCodeEnabled, - PasswordResetEnabled: settings.PasswordResetEnabled, - FrontendURL: settings.FrontendURL, - InvitationCodeEnabled: settings.InvitationCodeEnabled, - TotpEnabled: settings.TotpEnabled, - TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(), - LoginAgreementEnabled: settings.LoginAgreementEnabled, - LoginAgreementMode: settings.LoginAgreementMode, - LoginAgreementUpdatedAt: settings.LoginAgreementUpdatedAt, - LoginAgreementDocuments: loginAgreementDocumentsToDTO(settings.LoginAgreementDocuments), - SMTPHost: settings.SMTPHost, - SMTPPort: settings.SMTPPort, - SMTPUsername: settings.SMTPUsername, - SMTPPasswordConfigured: settings.SMTPPasswordConfigured, - SMTPFrom: settings.SMTPFrom, - SMTPFromName: settings.SMTPFromName, - SMTPUseTLS: settings.SMTPUseTLS, - TurnstileEnabled: settings.TurnstileEnabled, - TurnstileSiteKey: settings.TurnstileSiteKey, - TurnstileSecretKeyConfigured: settings.TurnstileSecretKeyConfigured, - APIKeyACLTrustForwardedIP: settings.APIKeyACLTrustForwardedIP, - LinuxDoConnectEnabled: settings.LinuxDoConnectEnabled, - LinuxDoConnectClientID: settings.LinuxDoConnectClientID, - LinuxDoConnectClientSecretConfigured: settings.LinuxDoConnectClientSecretConfigured, - LinuxDoConnectRedirectURL: settings.LinuxDoConnectRedirectURL, - DingTalkConnectEnabled: settings.DingTalkConnectEnabled, - DingTalkConnectClientID: settings.DingTalkConnectClientID, - DingTalkConnectClientSecretConfigured: settings.DingTalkConnectClientSecretConfigured, - DingTalkConnectRedirectURL: settings.DingTalkConnectRedirectURL, - DingTalkConnectCorpRestrictionPolicy: settings.DingTalkConnectCorpRestrictionPolicy, - DingTalkConnectInternalCorpID: settings.DingTalkConnectInternalCorpID, - DingTalkConnectBypassRegistration: settings.DingTalkConnectBypassRegistration, - DingTalkConnectSyncCorpEmail: settings.DingTalkConnectSyncCorpEmail, - DingTalkConnectSyncDisplayName: settings.DingTalkConnectSyncDisplayName, - DingTalkConnectSyncDept: settings.DingTalkConnectSyncDept, - DingTalkConnectSyncCorpEmailAttrKey: settings.DingTalkConnectSyncCorpEmailAttrKey, - DingTalkConnectSyncDisplayNameAttrKey: settings.DingTalkConnectSyncDisplayNameAttrKey, - DingTalkConnectSyncDeptAttrKey: settings.DingTalkConnectSyncDeptAttrKey, - DingTalkConnectSyncCorpEmailAttrName: settings.DingTalkConnectSyncCorpEmailAttrName, - DingTalkConnectSyncDisplayNameAttrName: settings.DingTalkConnectSyncDisplayNameAttrName, - DingTalkConnectSyncDeptAttrName: settings.DingTalkConnectSyncDeptAttrName, - WeChatConnectEnabled: settings.WeChatConnectEnabled, - WeChatConnectAppID: settings.WeChatConnectAppID, - WeChatConnectAppSecretConfigured: settings.WeChatConnectAppSecretConfigured, - WeChatConnectOpenAppID: settings.WeChatConnectOpenAppID, - WeChatConnectOpenAppSecretConfigured: settings.WeChatConnectOpenAppSecretConfigured, - WeChatConnectMPAppID: settings.WeChatConnectMPAppID, - WeChatConnectMPAppSecretConfigured: settings.WeChatConnectMPAppSecretConfigured, - WeChatConnectMobileAppID: settings.WeChatConnectMobileAppID, - WeChatConnectMobileAppSecretConfigured: settings.WeChatConnectMobileAppSecretConfigured, - WeChatConnectOpenEnabled: settings.WeChatConnectOpenEnabled, - WeChatConnectMPEnabled: settings.WeChatConnectMPEnabled, - WeChatConnectMobileEnabled: settings.WeChatConnectMobileEnabled, - WeChatConnectMode: settings.WeChatConnectMode, - WeChatConnectScopes: settings.WeChatConnectScopes, - WeChatConnectRedirectURL: settings.WeChatConnectRedirectURL, - WeChatConnectFrontendRedirectURL: settings.WeChatConnectFrontendRedirectURL, - OIDCConnectEnabled: settings.OIDCConnectEnabled, - OIDCConnectProviderName: settings.OIDCConnectProviderName, - OIDCConnectClientID: settings.OIDCConnectClientID, - OIDCConnectClientSecretConfigured: settings.OIDCConnectClientSecretConfigured, - OIDCConnectIssuerURL: settings.OIDCConnectIssuerURL, - OIDCConnectDiscoveryURL: settings.OIDCConnectDiscoveryURL, - OIDCConnectAuthorizeURL: settings.OIDCConnectAuthorizeURL, - OIDCConnectTokenURL: settings.OIDCConnectTokenURL, - OIDCConnectUserInfoURL: settings.OIDCConnectUserInfoURL, - OIDCConnectJWKSURL: settings.OIDCConnectJWKSURL, - OIDCConnectScopes: settings.OIDCConnectScopes, - OIDCConnectRedirectURL: settings.OIDCConnectRedirectURL, - OIDCConnectFrontendRedirectURL: settings.OIDCConnectFrontendRedirectURL, - OIDCConnectTokenAuthMethod: settings.OIDCConnectTokenAuthMethod, - OIDCConnectUsePKCE: settings.OIDCConnectUsePKCE, - OIDCConnectValidateIDToken: settings.OIDCConnectValidateIDToken, - OIDCConnectAllowedSigningAlgs: settings.OIDCConnectAllowedSigningAlgs, - OIDCConnectClockSkewSeconds: settings.OIDCConnectClockSkewSeconds, - OIDCConnectRequireEmailVerified: settings.OIDCConnectRequireEmailVerified, - OIDCConnectUserInfoEmailPath: settings.OIDCConnectUserInfoEmailPath, - OIDCConnectUserInfoIDPath: settings.OIDCConnectUserInfoIDPath, - OIDCConnectUserInfoUsernamePath: settings.OIDCConnectUserInfoUsernamePath, - GitHubOAuthEnabled: settings.GitHubOAuthEnabled, - GitHubOAuthClientID: settings.GitHubOAuthClientID, - GitHubOAuthClientSecretConfigured: settings.GitHubOAuthClientSecretConfigured, - GitHubOAuthRedirectURL: settings.GitHubOAuthRedirectURL, - GitHubOAuthFrontendRedirectURL: settings.GitHubOAuthFrontendRedirectURL, - GoogleOAuthEnabled: settings.GoogleOAuthEnabled, - GoogleOAuthClientID: settings.GoogleOAuthClientID, - GoogleOAuthClientSecretConfigured: settings.GoogleOAuthClientSecretConfigured, - GoogleOAuthRedirectURL: settings.GoogleOAuthRedirectURL, - GoogleOAuthFrontendRedirectURL: settings.GoogleOAuthFrontendRedirectURL, - SiteName: settings.SiteName, - SiteLogo: settings.SiteLogo, - SiteSubtitle: settings.SiteSubtitle, - APIBaseURL: settings.APIBaseURL, - ContactInfo: settings.ContactInfo, - DocURL: settings.DocURL, - HomeContent: settings.HomeContent, - HideCcsImportButton: settings.HideCcsImportButton, - PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled, - PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL, - TableDefaultPageSize: settings.TableDefaultPageSize, - TablePageSizeOptions: settings.TablePageSizeOptions, - CustomMenuItems: dto.ParseCustomMenuItems(settings.CustomMenuItems), - CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints), - DefaultConcurrency: settings.DefaultConcurrency, - DefaultBalance: settings.DefaultBalance, - RiskControlEnabled: settings.RiskControlEnabled, - CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled, - CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds, - AffiliateRebateRate: settings.AffiliateRebateRate, - AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours, - AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays, - AffiliateRebatePerInviteeCap: settings.AffiliateRebatePerInviteeCap, - DefaultUserRPMLimit: settings.DefaultUserRPMLimit, - DefaultSubscriptions: defaultSubscriptions, - EnableModelFallback: settings.EnableModelFallback, - FallbackModelAnthropic: settings.FallbackModelAnthropic, - FallbackModelOpenAI: settings.FallbackModelOpenAI, - FallbackModelGemini: settings.FallbackModelGemini, - FallbackModelAntigravity: settings.FallbackModelAntigravity, - EnableIdentityPatch: settings.EnableIdentityPatch, - IdentityPatchPrompt: settings.IdentityPatchPrompt, - OpsMonitoringEnabled: opsEnabled && settings.OpsMonitoringEnabled, - OpsRealtimeMonitoringEnabled: settings.OpsRealtimeMonitoringEnabled, - OpsQueryModeDefault: settings.OpsQueryModeDefault, - OpsMetricsIntervalSeconds: settings.OpsMetricsIntervalSeconds, - MinClaudeCodeVersion: settings.MinClaudeCodeVersion, - MaxClaudeCodeVersion: settings.MaxClaudeCodeVersion, - AllowUngroupedKeyScheduling: settings.AllowUngroupedKeyScheduling, - BackendModeEnabled: settings.BackendModeEnabled, - EnableFingerprintUnification: settings.EnableFingerprintUnification, - EnableMetadataPassthrough: settings.EnableMetadataPassthrough, - EnableCCHSigning: settings.EnableCCHSigning, - EnableClaudeOAuthSystemPromptInjection: settings.EnableClaudeOAuthSystemPromptInjection, - ClaudeOAuthSystemPrompt: settings.ClaudeOAuthSystemPrompt, - ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks, - EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection, - RewriteMessageCacheControl: settings.RewriteMessageCacheControl, - EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization, - AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion, - OpenAICodexUserAgent: settings.OpenAICodexUserAgent, - MinCodexVersion: settings.MinCodexVersion, - MaxCodexVersion: settings.MaxCodexVersion, - CodexCLIOnlyBlacklist: settings.CodexCLIOnlyBlacklist, - CodexCLIOnlyWhitelist: settings.CodexCLIOnlyWhitelist, - CodexCLIOnlyAllowAppServerClients: settings.CodexCLIOnlyAllowAppServerClients, - CodexCLIOnlyEngineFingerprintSignals: settings.CodexCLIOnlyEngineFingerprintSignals, - WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled, - PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource, - PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource, - PaymentVisibleMethodAlipayEnabled: settings.PaymentVisibleMethodAlipayEnabled, - PaymentVisibleMethodWxpayEnabled: settings.PaymentVisibleMethodWxpayEnabled, - OpenAIAdvancedSchedulerEnabled: settings.OpenAIAdvancedSchedulerEnabled, - BalanceLowNotifyEnabled: settings.BalanceLowNotifyEnabled, - BalanceLowNotifyThreshold: settings.BalanceLowNotifyThreshold, - BalanceLowNotifyRechargeURL: settings.BalanceLowNotifyRechargeURL, - SubscriptionExpiryNotifyEnabled: settings.SubscriptionExpiryNotifyEnabled, - AccountQuotaNotifyEnabled: settings.AccountQuotaNotifyEnabled, - AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(settings.AccountQuotaNotifyEmails), - PaymentEnabled: paymentCfg.Enabled, - PaymentMinAmount: paymentCfg.MinAmount, - PaymentMaxAmount: paymentCfg.MaxAmount, - PaymentDailyLimit: paymentCfg.DailyLimit, - PaymentOrderTimeoutMin: paymentCfg.OrderTimeoutMin, - PaymentMaxPendingOrders: paymentCfg.MaxPendingOrders, - PaymentEnabledTypes: paymentCfg.EnabledTypes, - PaymentBalanceDisabled: paymentCfg.BalanceDisabled, - PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier, - PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate, - PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy, - PaymentProductNamePrefix: paymentCfg.ProductNamePrefix, - PaymentProductNameSuffix: paymentCfg.ProductNameSuffix, - PaymentHelpImageURL: paymentCfg.HelpImageURL, - PaymentHelpText: paymentCfg.HelpText, - PaymentCancelRateLimitEnabled: paymentCfg.CancelRateLimitEnabled, - PaymentCancelRateLimitMax: paymentCfg.CancelRateLimitMax, - PaymentCancelRateLimitWindow: paymentCfg.CancelRateLimitWindow, - PaymentCancelRateLimitUnit: paymentCfg.CancelRateLimitUnit, - PaymentCancelRateLimitMode: paymentCfg.CancelRateLimitMode, - PaymentAlipayForceQRCode: paymentCfg.AlipayForceQRCode, + RegistrationEnabled: settings.RegistrationEnabled, + EmailVerifyEnabled: settings.EmailVerifyEnabled, + RegistrationEmailSuffixWhitelist: settings.RegistrationEmailSuffixWhitelist, + PromoCodeEnabled: settings.PromoCodeEnabled, + PasswordResetEnabled: settings.PasswordResetEnabled, + FrontendURL: settings.FrontendURL, + InvitationCodeEnabled: settings.InvitationCodeEnabled, + TotpEnabled: settings.TotpEnabled, + TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(), + LoginAgreementEnabled: settings.LoginAgreementEnabled, + LoginAgreementMode: settings.LoginAgreementMode, + LoginAgreementUpdatedAt: settings.LoginAgreementUpdatedAt, + LoginAgreementDocuments: loginAgreementDocumentsToDTO(settings.LoginAgreementDocuments), + SMTPHost: settings.SMTPHost, + SMTPPort: settings.SMTPPort, + SMTPUsername: settings.SMTPUsername, + SMTPPasswordConfigured: settings.SMTPPasswordConfigured, + SMTPFrom: settings.SMTPFrom, + SMTPFromName: settings.SMTPFromName, + SMTPUseTLS: settings.SMTPUseTLS, + TurnstileEnabled: settings.TurnstileEnabled, + TurnstileSiteKey: settings.TurnstileSiteKey, + TurnstileSecretKeyConfigured: settings.TurnstileSecretKeyConfigured, + APIKeyACLTrustForwardedIP: settings.APIKeyACLTrustForwardedIP, + LinuxDoConnectEnabled: settings.LinuxDoConnectEnabled, + LinuxDoConnectClientID: settings.LinuxDoConnectClientID, + LinuxDoConnectClientSecretConfigured: settings.LinuxDoConnectClientSecretConfigured, + LinuxDoConnectRedirectURL: settings.LinuxDoConnectRedirectURL, + DingTalkConnectEnabled: settings.DingTalkConnectEnabled, + DingTalkConnectClientID: settings.DingTalkConnectClientID, + DingTalkConnectClientSecretConfigured: settings.DingTalkConnectClientSecretConfigured, + DingTalkConnectRedirectURL: settings.DingTalkConnectRedirectURL, + DingTalkConnectCorpRestrictionPolicy: settings.DingTalkConnectCorpRestrictionPolicy, + DingTalkConnectInternalCorpID: settings.DingTalkConnectInternalCorpID, + DingTalkConnectBypassRegistration: settings.DingTalkConnectBypassRegistration, + DingTalkConnectSyncCorpEmail: settings.DingTalkConnectSyncCorpEmail, + DingTalkConnectSyncDisplayName: settings.DingTalkConnectSyncDisplayName, + DingTalkConnectSyncDept: settings.DingTalkConnectSyncDept, + DingTalkConnectSyncCorpEmailAttrKey: settings.DingTalkConnectSyncCorpEmailAttrKey, + DingTalkConnectSyncDisplayNameAttrKey: settings.DingTalkConnectSyncDisplayNameAttrKey, + DingTalkConnectSyncDeptAttrKey: settings.DingTalkConnectSyncDeptAttrKey, + DingTalkConnectSyncCorpEmailAttrName: settings.DingTalkConnectSyncCorpEmailAttrName, + DingTalkConnectSyncDisplayNameAttrName: settings.DingTalkConnectSyncDisplayNameAttrName, + DingTalkConnectSyncDeptAttrName: settings.DingTalkConnectSyncDeptAttrName, + WeChatConnectEnabled: settings.WeChatConnectEnabled, + WeChatConnectAppID: settings.WeChatConnectAppID, + WeChatConnectAppSecretConfigured: settings.WeChatConnectAppSecretConfigured, + WeChatConnectOpenAppID: settings.WeChatConnectOpenAppID, + WeChatConnectOpenAppSecretConfigured: settings.WeChatConnectOpenAppSecretConfigured, + WeChatConnectMPAppID: settings.WeChatConnectMPAppID, + WeChatConnectMPAppSecretConfigured: settings.WeChatConnectMPAppSecretConfigured, + WeChatConnectMobileAppID: settings.WeChatConnectMobileAppID, + WeChatConnectMobileAppSecretConfigured: settings.WeChatConnectMobileAppSecretConfigured, + WeChatConnectOpenEnabled: settings.WeChatConnectOpenEnabled, + WeChatConnectMPEnabled: settings.WeChatConnectMPEnabled, + WeChatConnectMobileEnabled: settings.WeChatConnectMobileEnabled, + WeChatConnectMode: settings.WeChatConnectMode, + WeChatConnectScopes: settings.WeChatConnectScopes, + WeChatConnectRedirectURL: settings.WeChatConnectRedirectURL, + WeChatConnectFrontendRedirectURL: settings.WeChatConnectFrontendRedirectURL, + OIDCConnectEnabled: settings.OIDCConnectEnabled, + OIDCConnectProviderName: settings.OIDCConnectProviderName, + OIDCConnectClientID: settings.OIDCConnectClientID, + OIDCConnectClientSecretConfigured: settings.OIDCConnectClientSecretConfigured, + OIDCConnectIssuerURL: settings.OIDCConnectIssuerURL, + OIDCConnectDiscoveryURL: settings.OIDCConnectDiscoveryURL, + OIDCConnectAuthorizeURL: settings.OIDCConnectAuthorizeURL, + OIDCConnectTokenURL: settings.OIDCConnectTokenURL, + OIDCConnectUserInfoURL: settings.OIDCConnectUserInfoURL, + OIDCConnectJWKSURL: settings.OIDCConnectJWKSURL, + OIDCConnectScopes: settings.OIDCConnectScopes, + OIDCConnectRedirectURL: settings.OIDCConnectRedirectURL, + OIDCConnectFrontendRedirectURL: settings.OIDCConnectFrontendRedirectURL, + OIDCConnectTokenAuthMethod: settings.OIDCConnectTokenAuthMethod, + OIDCConnectUsePKCE: settings.OIDCConnectUsePKCE, + OIDCConnectValidateIDToken: settings.OIDCConnectValidateIDToken, + OIDCConnectAllowedSigningAlgs: settings.OIDCConnectAllowedSigningAlgs, + OIDCConnectClockSkewSeconds: settings.OIDCConnectClockSkewSeconds, + OIDCConnectRequireEmailVerified: settings.OIDCConnectRequireEmailVerified, + OIDCConnectUserInfoEmailPath: settings.OIDCConnectUserInfoEmailPath, + OIDCConnectUserInfoIDPath: settings.OIDCConnectUserInfoIDPath, + OIDCConnectUserInfoUsernamePath: settings.OIDCConnectUserInfoUsernamePath, + GitHubOAuthEnabled: settings.GitHubOAuthEnabled, + GitHubOAuthClientID: settings.GitHubOAuthClientID, + GitHubOAuthClientSecretConfigured: settings.GitHubOAuthClientSecretConfigured, + GitHubOAuthRedirectURL: settings.GitHubOAuthRedirectURL, + GitHubOAuthFrontendRedirectURL: settings.GitHubOAuthFrontendRedirectURL, + GoogleOAuthEnabled: settings.GoogleOAuthEnabled, + GoogleOAuthClientID: settings.GoogleOAuthClientID, + GoogleOAuthClientSecretConfigured: settings.GoogleOAuthClientSecretConfigured, + GoogleOAuthRedirectURL: settings.GoogleOAuthRedirectURL, + GoogleOAuthFrontendRedirectURL: settings.GoogleOAuthFrontendRedirectURL, + SiteName: settings.SiteName, + SiteLogo: settings.SiteLogo, + SiteSubtitle: settings.SiteSubtitle, + APIBaseURL: settings.APIBaseURL, + ContactInfo: settings.ContactInfo, + DocURL: settings.DocURL, + HomeContent: settings.HomeContent, + HideCcsImportButton: settings.HideCcsImportButton, + PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled, + PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL, + TableDefaultPageSize: settings.TableDefaultPageSize, + TablePageSizeOptions: settings.TablePageSizeOptions, + CustomMenuItems: dto.ParseCustomMenuItems(settings.CustomMenuItems), + CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints), + DefaultConcurrency: settings.DefaultConcurrency, + DefaultBalance: settings.DefaultBalance, + RiskControlEnabled: settings.RiskControlEnabled, + CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled, + CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds, + AffiliateRebateRate: settings.AffiliateRebateRate, + AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours, + AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays, + AffiliateRebatePerInviteeCap: settings.AffiliateRebatePerInviteeCap, + DefaultUserRPMLimit: settings.DefaultUserRPMLimit, + DefaultSubscriptions: defaultSubscriptions, + EnableModelFallback: settings.EnableModelFallback, + FallbackModelAnthropic: settings.FallbackModelAnthropic, + FallbackModelOpenAI: settings.FallbackModelOpenAI, + FallbackModelGemini: settings.FallbackModelGemini, + FallbackModelAntigravity: settings.FallbackModelAntigravity, + EnableIdentityPatch: settings.EnableIdentityPatch, + IdentityPatchPrompt: settings.IdentityPatchPrompt, + OpsMonitoringEnabled: opsEnabled && settings.OpsMonitoringEnabled, + OpsRealtimeMonitoringEnabled: settings.OpsRealtimeMonitoringEnabled, + OpsQueryModeDefault: settings.OpsQueryModeDefault, + OpsMetricsIntervalSeconds: settings.OpsMetricsIntervalSeconds, + MinClaudeCodeVersion: settings.MinClaudeCodeVersion, + MaxClaudeCodeVersion: settings.MaxClaudeCodeVersion, + AllowUngroupedKeyScheduling: settings.AllowUngroupedKeyScheduling, + BackendModeEnabled: settings.BackendModeEnabled, + EnableFingerprintUnification: settings.EnableFingerprintUnification, + EnableMetadataPassthrough: settings.EnableMetadataPassthrough, + EnableCCHSigning: settings.EnableCCHSigning, + EnableClaudeOAuthSystemPromptInjection: settings.EnableClaudeOAuthSystemPromptInjection, + ClaudeOAuthSystemPrompt: settings.ClaudeOAuthSystemPrompt, + ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks, + EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection, + RewriteMessageCacheControl: settings.RewriteMessageCacheControl, + EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization, + AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion, + OpenAICodexUserAgent: settings.OpenAICodexUserAgent, + MinCodexVersion: settings.MinCodexVersion, + MaxCodexVersion: settings.MaxCodexVersion, + CodexCLIOnlyBlacklist: settings.CodexCLIOnlyBlacklist, + CodexCLIOnlyWhitelist: settings.CodexCLIOnlyWhitelist, + CodexCLIOnlyAllowAppServerClients: settings.CodexCLIOnlyAllowAppServerClients, + CodexCLIOnlyEngineFingerprintSignals: settings.CodexCLIOnlyEngineFingerprintSignals, + WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled, + PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource, + PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource, + PaymentVisibleMethodAlipayEnabled: settings.PaymentVisibleMethodAlipayEnabled, + PaymentVisibleMethodWxpayEnabled: settings.PaymentVisibleMethodWxpayEnabled, + OpenAIAdvancedSchedulerEnabled: settings.OpenAIAdvancedSchedulerEnabled, + OpenAIAdvancedSchedulerStickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled, + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled, + OpenAIAdvancedSchedulerLBTopK: settings.OpenAIAdvancedSchedulerLBTopK, + OpenAIAdvancedSchedulerWeightPriority: settings.OpenAIAdvancedSchedulerWeightPriority, + OpenAIAdvancedSchedulerWeightLoad: settings.OpenAIAdvancedSchedulerWeightLoad, + OpenAIAdvancedSchedulerWeightQueue: settings.OpenAIAdvancedSchedulerWeightQueue, + OpenAIAdvancedSchedulerWeightErrorRate: settings.OpenAIAdvancedSchedulerWeightErrorRate, + OpenAIAdvancedSchedulerWeightTTFT: settings.OpenAIAdvancedSchedulerWeightTTFT, + OpenAIAdvancedSchedulerWeightReset: settings.OpenAIAdvancedSchedulerWeightReset, + OpenAIAdvancedSchedulerWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, + OpenAIAdvancedSchedulerWeightPreviousResponse: settings.OpenAIAdvancedSchedulerWeightPreviousResponse, + OpenAIAdvancedSchedulerWeightSessionSticky: settings.OpenAIAdvancedSchedulerWeightSessionSticky, + OpenAIAdvancedSchedulerEffectiveLBTopK: settings.OpenAIAdvancedSchedulerEffectiveLBTopK, + OpenAIAdvancedSchedulerEffectiveWeightPriority: settings.OpenAIAdvancedSchedulerEffectiveWeightPriority, + OpenAIAdvancedSchedulerEffectiveWeightLoad: settings.OpenAIAdvancedSchedulerEffectiveWeightLoad, + OpenAIAdvancedSchedulerEffectiveWeightQueue: settings.OpenAIAdvancedSchedulerEffectiveWeightQueue, + OpenAIAdvancedSchedulerEffectiveWeightErrorRate: settings.OpenAIAdvancedSchedulerEffectiveWeightErrorRate, + OpenAIAdvancedSchedulerEffectiveWeightTTFT: settings.OpenAIAdvancedSchedulerEffectiveWeightTTFT, + OpenAIAdvancedSchedulerEffectiveWeightReset: settings.OpenAIAdvancedSchedulerEffectiveWeightReset, + OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom, + OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: settings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse, + OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky, + BalanceLowNotifyEnabled: settings.BalanceLowNotifyEnabled, + BalanceLowNotifyThreshold: settings.BalanceLowNotifyThreshold, + BalanceLowNotifyRechargeURL: settings.BalanceLowNotifyRechargeURL, + SubscriptionExpiryNotifyEnabled: settings.SubscriptionExpiryNotifyEnabled, + AccountQuotaNotifyEnabled: settings.AccountQuotaNotifyEnabled, + AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(settings.AccountQuotaNotifyEmails), + PaymentEnabled: paymentCfg.Enabled, + PaymentMinAmount: paymentCfg.MinAmount, + PaymentMaxAmount: paymentCfg.MaxAmount, + PaymentDailyLimit: paymentCfg.DailyLimit, + PaymentOrderTimeoutMin: paymentCfg.OrderTimeoutMin, + PaymentMaxPendingOrders: paymentCfg.MaxPendingOrders, + PaymentEnabledTypes: paymentCfg.EnabledTypes, + PaymentBalanceDisabled: paymentCfg.BalanceDisabled, + PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier, + PaymentSubscriptionUSDToCNYRate: paymentCfg.SubscriptionUSDToCNYRate, + PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate, + PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy, + PaymentProductNamePrefix: paymentCfg.ProductNamePrefix, + PaymentProductNameSuffix: paymentCfg.ProductNameSuffix, + PaymentHelpImageURL: paymentCfg.HelpImageURL, + PaymentHelpText: paymentCfg.HelpText, + PaymentCancelRateLimitEnabled: paymentCfg.CancelRateLimitEnabled, + PaymentCancelRateLimitMax: paymentCfg.CancelRateLimitMax, + PaymentCancelRateLimitWindow: paymentCfg.CancelRateLimitWindow, + PaymentCancelRateLimitUnit: paymentCfg.CancelRateLimitUnit, + PaymentCancelRateLimitMode: paymentCfg.CancelRateLimitMode, + PaymentAlipayForceQRCode: paymentCfg.AlipayForceQRCode, ChannelMonitorEnabled: settings.ChannelMonitorEnabled, ChannelMonitorDefaultIntervalSeconds: settings.ChannelMonitorDefaultIntervalSeconds, @@ -618,7 +641,19 @@ type UpdateSettingsRequest struct { PaymentVisibleMethodWxpayEnabled *bool `json:"payment_visible_method_wxpay_enabled"` // OpenAI account scheduling - OpenAIAdvancedSchedulerEnabled *bool `json:"openai_advanced_scheduler_enabled"` + OpenAIAdvancedSchedulerEnabled *bool `json:"openai_advanced_scheduler_enabled"` + OpenAIAdvancedSchedulerStickyWeightedEnabled *bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"` + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled *bool `json:"openai_advanced_scheduler_subscription_priority_enabled"` + OpenAIAdvancedSchedulerLBTopK *string `json:"openai_advanced_scheduler_lb_top_k"` + OpenAIAdvancedSchedulerWeightPriority *string `json:"openai_advanced_scheduler_weight_priority"` + OpenAIAdvancedSchedulerWeightLoad *string `json:"openai_advanced_scheduler_weight_load"` + OpenAIAdvancedSchedulerWeightQueue *string `json:"openai_advanced_scheduler_weight_queue"` + OpenAIAdvancedSchedulerWeightErrorRate *string `json:"openai_advanced_scheduler_weight_error_rate"` + OpenAIAdvancedSchedulerWeightTTFT *string `json:"openai_advanced_scheduler_weight_ttft"` + OpenAIAdvancedSchedulerWeightReset *string `json:"openai_advanced_scheduler_weight_reset"` + OpenAIAdvancedSchedulerWeightQuotaHeadroom *string `json:"openai_advanced_scheduler_weight_quota_headroom"` + OpenAIAdvancedSchedulerWeightPreviousResponse *string `json:"openai_advanced_scheduler_weight_previous_response"` + OpenAIAdvancedSchedulerWeightSessionSticky *string `json:"openai_advanced_scheduler_weight_session_sticky"` // 余额不足提醒 BalanceLowNotifyEnabled *bool `json:"balance_low_notify_enabled"` @@ -638,6 +673,7 @@ type UpdateSettingsRequest struct { PaymentEnabledTypes []string `json:"payment_enabled_types"` PaymentBalanceDisabled *bool `json:"payment_balance_disabled"` PaymentBalanceRechargeMultiplier *float64 `json:"payment_balance_recharge_multiplier"` + PaymentSubscriptionUSDToCNYRate *float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate *float64 `json:"payment_recharge_fee_rate"` PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` @@ -1792,6 +1828,28 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.OpenAIAdvancedSchedulerEnabled }(), + OpenAIAdvancedSchedulerStickyWeightedEnabled: func() bool { + if req.OpenAIAdvancedSchedulerStickyWeightedEnabled != nil { + return *req.OpenAIAdvancedSchedulerStickyWeightedEnabled + } + return previousSettings.OpenAIAdvancedSchedulerStickyWeightedEnabled + }(), + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: func() bool { + if req.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled != nil { + return *req.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled + } + return previousSettings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled + }(), + OpenAIAdvancedSchedulerLBTopK: stringSetting(req.OpenAIAdvancedSchedulerLBTopK, previousSettings.OpenAIAdvancedSchedulerLBTopK), + OpenAIAdvancedSchedulerWeightPriority: stringSetting(req.OpenAIAdvancedSchedulerWeightPriority, previousSettings.OpenAIAdvancedSchedulerWeightPriority), + OpenAIAdvancedSchedulerWeightLoad: stringSetting(req.OpenAIAdvancedSchedulerWeightLoad, previousSettings.OpenAIAdvancedSchedulerWeightLoad), + OpenAIAdvancedSchedulerWeightQueue: stringSetting(req.OpenAIAdvancedSchedulerWeightQueue, previousSettings.OpenAIAdvancedSchedulerWeightQueue), + OpenAIAdvancedSchedulerWeightErrorRate: stringSetting(req.OpenAIAdvancedSchedulerWeightErrorRate, previousSettings.OpenAIAdvancedSchedulerWeightErrorRate), + OpenAIAdvancedSchedulerWeightTTFT: stringSetting(req.OpenAIAdvancedSchedulerWeightTTFT, previousSettings.OpenAIAdvancedSchedulerWeightTTFT), + OpenAIAdvancedSchedulerWeightReset: stringSetting(req.OpenAIAdvancedSchedulerWeightReset, previousSettings.OpenAIAdvancedSchedulerWeightReset), + OpenAIAdvancedSchedulerWeightQuotaHeadroom: stringSetting(req.OpenAIAdvancedSchedulerWeightQuotaHeadroom, previousSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom), + OpenAIAdvancedSchedulerWeightPreviousResponse: stringSetting(req.OpenAIAdvancedSchedulerWeightPreviousResponse, previousSettings.OpenAIAdvancedSchedulerWeightPreviousResponse), + OpenAIAdvancedSchedulerWeightSessionSticky: stringSetting(req.OpenAIAdvancedSchedulerWeightSessionSticky, previousSettings.OpenAIAdvancedSchedulerWeightSessionSticky), BalanceLowNotifyEnabled: func() bool { if req.BalanceLowNotifyEnabled != nil { return *req.BalanceLowNotifyEnabled @@ -1959,6 +2017,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { EnabledTypes: req.PaymentEnabledTypes, BalanceDisabled: req.PaymentBalanceDisabled, BalanceRechargeMultiplier: req.PaymentBalanceRechargeMultiplier, + SubscriptionUSDToCNYRate: req.PaymentSubscriptionUSDToCNYRate, RechargeFeeRate: req.PaymentRechargeFeeRate, LoadBalanceStrategy: req.PaymentLoadBalanceStrat, ProductNamePrefix: req.PaymentProductNamePrefix, @@ -2014,184 +2073,207 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } payload := dto.SystemSettings{ - RegistrationEnabled: updatedSettings.RegistrationEnabled, - EmailVerifyEnabled: updatedSettings.EmailVerifyEnabled, - RegistrationEmailSuffixWhitelist: updatedSettings.RegistrationEmailSuffixWhitelist, - PromoCodeEnabled: updatedSettings.PromoCodeEnabled, - PasswordResetEnabled: updatedSettings.PasswordResetEnabled, - FrontendURL: updatedSettings.FrontendURL, - InvitationCodeEnabled: updatedSettings.InvitationCodeEnabled, - TotpEnabled: updatedSettings.TotpEnabled, - TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(), - LoginAgreementEnabled: updatedSettings.LoginAgreementEnabled, - LoginAgreementMode: updatedSettings.LoginAgreementMode, - LoginAgreementUpdatedAt: updatedSettings.LoginAgreementUpdatedAt, - LoginAgreementDocuments: loginAgreementDocumentsToDTO(updatedSettings.LoginAgreementDocuments), - SMTPHost: updatedSettings.SMTPHost, - SMTPPort: updatedSettings.SMTPPort, - SMTPUsername: updatedSettings.SMTPUsername, - SMTPPasswordConfigured: updatedSettings.SMTPPasswordConfigured, - SMTPFrom: updatedSettings.SMTPFrom, - SMTPFromName: updatedSettings.SMTPFromName, - SMTPUseTLS: updatedSettings.SMTPUseTLS, - TurnstileEnabled: updatedSettings.TurnstileEnabled, - TurnstileSiteKey: updatedSettings.TurnstileSiteKey, - TurnstileSecretKeyConfigured: updatedSettings.TurnstileSecretKeyConfigured, - APIKeyACLTrustForwardedIP: updatedSettings.APIKeyACLTrustForwardedIP, - LinuxDoConnectEnabled: updatedSettings.LinuxDoConnectEnabled, - LinuxDoConnectClientID: updatedSettings.LinuxDoConnectClientID, - LinuxDoConnectClientSecretConfigured: updatedSettings.LinuxDoConnectClientSecretConfigured, - LinuxDoConnectRedirectURL: updatedSettings.LinuxDoConnectRedirectURL, - DingTalkConnectEnabled: updatedSettings.DingTalkConnectEnabled, - DingTalkConnectClientID: updatedSettings.DingTalkConnectClientID, - DingTalkConnectClientSecretConfigured: updatedSettings.DingTalkConnectClientSecretConfigured, - DingTalkConnectRedirectURL: updatedSettings.DingTalkConnectRedirectURL, - DingTalkConnectCorpRestrictionPolicy: updatedSettings.DingTalkConnectCorpRestrictionPolicy, - DingTalkConnectInternalCorpID: updatedSettings.DingTalkConnectInternalCorpID, - DingTalkConnectBypassRegistration: updatedSettings.DingTalkConnectBypassRegistration, - DingTalkConnectSyncCorpEmail: updatedSettings.DingTalkConnectSyncCorpEmail, - DingTalkConnectSyncDisplayName: updatedSettings.DingTalkConnectSyncDisplayName, - DingTalkConnectSyncDept: updatedSettings.DingTalkConnectSyncDept, - DingTalkConnectSyncCorpEmailAttrKey: updatedSettings.DingTalkConnectSyncCorpEmailAttrKey, - DingTalkConnectSyncDisplayNameAttrKey: updatedSettings.DingTalkConnectSyncDisplayNameAttrKey, - DingTalkConnectSyncDeptAttrKey: updatedSettings.DingTalkConnectSyncDeptAttrKey, - DingTalkConnectSyncCorpEmailAttrName: updatedSettings.DingTalkConnectSyncCorpEmailAttrName, - DingTalkConnectSyncDisplayNameAttrName: updatedSettings.DingTalkConnectSyncDisplayNameAttrName, - DingTalkConnectSyncDeptAttrName: updatedSettings.DingTalkConnectSyncDeptAttrName, - WeChatConnectEnabled: updatedSettings.WeChatConnectEnabled, - WeChatConnectAppID: updatedSettings.WeChatConnectAppID, - WeChatConnectAppSecretConfigured: updatedSettings.WeChatConnectAppSecretConfigured, - WeChatConnectOpenAppID: updatedSettings.WeChatConnectOpenAppID, - WeChatConnectOpenAppSecretConfigured: updatedSettings.WeChatConnectOpenAppSecretConfigured, - WeChatConnectMPAppID: updatedSettings.WeChatConnectMPAppID, - WeChatConnectMPAppSecretConfigured: updatedSettings.WeChatConnectMPAppSecretConfigured, - WeChatConnectMobileAppID: updatedSettings.WeChatConnectMobileAppID, - WeChatConnectMobileAppSecretConfigured: updatedSettings.WeChatConnectMobileAppSecretConfigured, - WeChatConnectOpenEnabled: updatedSettings.WeChatConnectOpenEnabled, - WeChatConnectMPEnabled: updatedSettings.WeChatConnectMPEnabled, - WeChatConnectMobileEnabled: updatedSettings.WeChatConnectMobileEnabled, - WeChatConnectMode: updatedSettings.WeChatConnectMode, - WeChatConnectScopes: updatedSettings.WeChatConnectScopes, - WeChatConnectRedirectURL: updatedSettings.WeChatConnectRedirectURL, - WeChatConnectFrontendRedirectURL: updatedSettings.WeChatConnectFrontendRedirectURL, - OIDCConnectEnabled: updatedSettings.OIDCConnectEnabled, - OIDCConnectProviderName: updatedSettings.OIDCConnectProviderName, - OIDCConnectClientID: updatedSettings.OIDCConnectClientID, - OIDCConnectClientSecretConfigured: updatedSettings.OIDCConnectClientSecretConfigured, - OIDCConnectIssuerURL: updatedSettings.OIDCConnectIssuerURL, - OIDCConnectDiscoveryURL: updatedSettings.OIDCConnectDiscoveryURL, - OIDCConnectAuthorizeURL: updatedSettings.OIDCConnectAuthorizeURL, - OIDCConnectTokenURL: updatedSettings.OIDCConnectTokenURL, - OIDCConnectUserInfoURL: updatedSettings.OIDCConnectUserInfoURL, - OIDCConnectJWKSURL: updatedSettings.OIDCConnectJWKSURL, - OIDCConnectScopes: updatedSettings.OIDCConnectScopes, - OIDCConnectRedirectURL: updatedSettings.OIDCConnectRedirectURL, - OIDCConnectFrontendRedirectURL: updatedSettings.OIDCConnectFrontendRedirectURL, - OIDCConnectTokenAuthMethod: updatedSettings.OIDCConnectTokenAuthMethod, - OIDCConnectUsePKCE: updatedSettings.OIDCConnectUsePKCE, - OIDCConnectValidateIDToken: updatedSettings.OIDCConnectValidateIDToken, - OIDCConnectAllowedSigningAlgs: updatedSettings.OIDCConnectAllowedSigningAlgs, - OIDCConnectClockSkewSeconds: updatedSettings.OIDCConnectClockSkewSeconds, - OIDCConnectRequireEmailVerified: updatedSettings.OIDCConnectRequireEmailVerified, - OIDCConnectUserInfoEmailPath: updatedSettings.OIDCConnectUserInfoEmailPath, - OIDCConnectUserInfoIDPath: updatedSettings.OIDCConnectUserInfoIDPath, - OIDCConnectUserInfoUsernamePath: updatedSettings.OIDCConnectUserInfoUsernamePath, - GitHubOAuthEnabled: updatedSettings.GitHubOAuthEnabled, - GitHubOAuthClientID: updatedSettings.GitHubOAuthClientID, - GitHubOAuthClientSecretConfigured: updatedSettings.GitHubOAuthClientSecretConfigured, - GitHubOAuthRedirectURL: updatedSettings.GitHubOAuthRedirectURL, - GitHubOAuthFrontendRedirectURL: updatedSettings.GitHubOAuthFrontendRedirectURL, - GoogleOAuthEnabled: updatedSettings.GoogleOAuthEnabled, - GoogleOAuthClientID: updatedSettings.GoogleOAuthClientID, - GoogleOAuthClientSecretConfigured: updatedSettings.GoogleOAuthClientSecretConfigured, - GoogleOAuthRedirectURL: updatedSettings.GoogleOAuthRedirectURL, - GoogleOAuthFrontendRedirectURL: updatedSettings.GoogleOAuthFrontendRedirectURL, - SiteName: updatedSettings.SiteName, - SiteLogo: updatedSettings.SiteLogo, - SiteSubtitle: updatedSettings.SiteSubtitle, - APIBaseURL: updatedSettings.APIBaseURL, - ContactInfo: updatedSettings.ContactInfo, - DocURL: updatedSettings.DocURL, - HomeContent: updatedSettings.HomeContent, - HideCcsImportButton: updatedSettings.HideCcsImportButton, - PurchaseSubscriptionEnabled: updatedSettings.PurchaseSubscriptionEnabled, - PurchaseSubscriptionURL: updatedSettings.PurchaseSubscriptionURL, - TableDefaultPageSize: updatedSettings.TableDefaultPageSize, - TablePageSizeOptions: updatedSettings.TablePageSizeOptions, - CustomMenuItems: dto.ParseCustomMenuItems(updatedSettings.CustomMenuItems), - CustomEndpoints: dto.ParseCustomEndpoints(updatedSettings.CustomEndpoints), - DefaultConcurrency: updatedSettings.DefaultConcurrency, - DefaultBalance: updatedSettings.DefaultBalance, - AffiliateRebateRate: updatedSettings.AffiliateRebateRate, - AffiliateRebateFreezeHours: updatedSettings.AffiliateRebateFreezeHours, - AffiliateRebateDurationDays: updatedSettings.AffiliateRebateDurationDays, - AffiliateRebatePerInviteeCap: updatedSettings.AffiliateRebatePerInviteeCap, - DefaultUserRPMLimit: updatedSettings.DefaultUserRPMLimit, - DefaultSubscriptions: updatedDefaultSubscriptions, - EnableModelFallback: updatedSettings.EnableModelFallback, - FallbackModelAnthropic: updatedSettings.FallbackModelAnthropic, - FallbackModelOpenAI: updatedSettings.FallbackModelOpenAI, - FallbackModelGemini: updatedSettings.FallbackModelGemini, - FallbackModelAntigravity: updatedSettings.FallbackModelAntigravity, - EnableIdentityPatch: updatedSettings.EnableIdentityPatch, - IdentityPatchPrompt: updatedSettings.IdentityPatchPrompt, - OpsMonitoringEnabled: updatedSettings.OpsMonitoringEnabled, - OpsRealtimeMonitoringEnabled: updatedSettings.OpsRealtimeMonitoringEnabled, - OpsQueryModeDefault: updatedSettings.OpsQueryModeDefault, - OpsMetricsIntervalSeconds: updatedSettings.OpsMetricsIntervalSeconds, - MinClaudeCodeVersion: updatedSettings.MinClaudeCodeVersion, - MaxClaudeCodeVersion: updatedSettings.MaxClaudeCodeVersion, - AllowUngroupedKeyScheduling: updatedSettings.AllowUngroupedKeyScheduling, - BackendModeEnabled: updatedSettings.BackendModeEnabled, - EnableFingerprintUnification: updatedSettings.EnableFingerprintUnification, - EnableMetadataPassthrough: updatedSettings.EnableMetadataPassthrough, - EnableCCHSigning: updatedSettings.EnableCCHSigning, - EnableClaudeOAuthSystemPromptInjection: updatedSettings.EnableClaudeOAuthSystemPromptInjection, - ClaudeOAuthSystemPrompt: updatedSettings.ClaudeOAuthSystemPrompt, - ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks, - EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection, - RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl, - EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization, - AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion, - OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent, - MinCodexVersion: updatedSettings.MinCodexVersion, - MaxCodexVersion: updatedSettings.MaxCodexVersion, - CodexCLIOnlyBlacklist: updatedSettings.CodexCLIOnlyBlacklist, - CodexCLIOnlyWhitelist: updatedSettings.CodexCLIOnlyWhitelist, - CodexCLIOnlyAllowAppServerClients: updatedSettings.CodexCLIOnlyAllowAppServerClients, - CodexCLIOnlyEngineFingerprintSignals: updatedSettings.CodexCLIOnlyEngineFingerprintSignals, - PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource, - PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource, - PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled, - PaymentVisibleMethodWxpayEnabled: updatedSettings.PaymentVisibleMethodWxpayEnabled, - OpenAIAdvancedSchedulerEnabled: updatedSettings.OpenAIAdvancedSchedulerEnabled, - BalanceLowNotifyEnabled: updatedSettings.BalanceLowNotifyEnabled, - BalanceLowNotifyThreshold: updatedSettings.BalanceLowNotifyThreshold, - BalanceLowNotifyRechargeURL: updatedSettings.BalanceLowNotifyRechargeURL, - SubscriptionExpiryNotifyEnabled: updatedSettings.SubscriptionExpiryNotifyEnabled, - AccountQuotaNotifyEnabled: updatedSettings.AccountQuotaNotifyEnabled, - AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(updatedSettings.AccountQuotaNotifyEmails), - PaymentEnabled: updatedPaymentCfg.Enabled, - PaymentMinAmount: updatedPaymentCfg.MinAmount, - PaymentMaxAmount: updatedPaymentCfg.MaxAmount, - PaymentDailyLimit: updatedPaymentCfg.DailyLimit, - PaymentOrderTimeoutMin: updatedPaymentCfg.OrderTimeoutMin, - PaymentMaxPendingOrders: updatedPaymentCfg.MaxPendingOrders, - PaymentEnabledTypes: updatedPaymentCfg.EnabledTypes, - PaymentBalanceDisabled: updatedPaymentCfg.BalanceDisabled, - PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier, - PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate, - PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy, - PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix, - PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix, - PaymentHelpImageURL: updatedPaymentCfg.HelpImageURL, - PaymentHelpText: updatedPaymentCfg.HelpText, - PaymentCancelRateLimitEnabled: updatedPaymentCfg.CancelRateLimitEnabled, - PaymentCancelRateLimitMax: updatedPaymentCfg.CancelRateLimitMax, - PaymentCancelRateLimitWindow: updatedPaymentCfg.CancelRateLimitWindow, - PaymentCancelRateLimitUnit: updatedPaymentCfg.CancelRateLimitUnit, - PaymentCancelRateLimitMode: updatedPaymentCfg.CancelRateLimitMode, - PaymentAlipayForceQRCode: updatedPaymentCfg.AlipayForceQRCode, + RegistrationEnabled: updatedSettings.RegistrationEnabled, + EmailVerifyEnabled: updatedSettings.EmailVerifyEnabled, + RegistrationEmailSuffixWhitelist: updatedSettings.RegistrationEmailSuffixWhitelist, + PromoCodeEnabled: updatedSettings.PromoCodeEnabled, + PasswordResetEnabled: updatedSettings.PasswordResetEnabled, + FrontendURL: updatedSettings.FrontendURL, + InvitationCodeEnabled: updatedSettings.InvitationCodeEnabled, + TotpEnabled: updatedSettings.TotpEnabled, + TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(), + LoginAgreementEnabled: updatedSettings.LoginAgreementEnabled, + LoginAgreementMode: updatedSettings.LoginAgreementMode, + LoginAgreementUpdatedAt: updatedSettings.LoginAgreementUpdatedAt, + LoginAgreementDocuments: loginAgreementDocumentsToDTO(updatedSettings.LoginAgreementDocuments), + SMTPHost: updatedSettings.SMTPHost, + SMTPPort: updatedSettings.SMTPPort, + SMTPUsername: updatedSettings.SMTPUsername, + SMTPPasswordConfigured: updatedSettings.SMTPPasswordConfigured, + SMTPFrom: updatedSettings.SMTPFrom, + SMTPFromName: updatedSettings.SMTPFromName, + SMTPUseTLS: updatedSettings.SMTPUseTLS, + TurnstileEnabled: updatedSettings.TurnstileEnabled, + TurnstileSiteKey: updatedSettings.TurnstileSiteKey, + TurnstileSecretKeyConfigured: updatedSettings.TurnstileSecretKeyConfigured, + APIKeyACLTrustForwardedIP: updatedSettings.APIKeyACLTrustForwardedIP, + LinuxDoConnectEnabled: updatedSettings.LinuxDoConnectEnabled, + LinuxDoConnectClientID: updatedSettings.LinuxDoConnectClientID, + LinuxDoConnectClientSecretConfigured: updatedSettings.LinuxDoConnectClientSecretConfigured, + LinuxDoConnectRedirectURL: updatedSettings.LinuxDoConnectRedirectURL, + DingTalkConnectEnabled: updatedSettings.DingTalkConnectEnabled, + DingTalkConnectClientID: updatedSettings.DingTalkConnectClientID, + DingTalkConnectClientSecretConfigured: updatedSettings.DingTalkConnectClientSecretConfigured, + DingTalkConnectRedirectURL: updatedSettings.DingTalkConnectRedirectURL, + DingTalkConnectCorpRestrictionPolicy: updatedSettings.DingTalkConnectCorpRestrictionPolicy, + DingTalkConnectInternalCorpID: updatedSettings.DingTalkConnectInternalCorpID, + DingTalkConnectBypassRegistration: updatedSettings.DingTalkConnectBypassRegistration, + DingTalkConnectSyncCorpEmail: updatedSettings.DingTalkConnectSyncCorpEmail, + DingTalkConnectSyncDisplayName: updatedSettings.DingTalkConnectSyncDisplayName, + DingTalkConnectSyncDept: updatedSettings.DingTalkConnectSyncDept, + DingTalkConnectSyncCorpEmailAttrKey: updatedSettings.DingTalkConnectSyncCorpEmailAttrKey, + DingTalkConnectSyncDisplayNameAttrKey: updatedSettings.DingTalkConnectSyncDisplayNameAttrKey, + DingTalkConnectSyncDeptAttrKey: updatedSettings.DingTalkConnectSyncDeptAttrKey, + DingTalkConnectSyncCorpEmailAttrName: updatedSettings.DingTalkConnectSyncCorpEmailAttrName, + DingTalkConnectSyncDisplayNameAttrName: updatedSettings.DingTalkConnectSyncDisplayNameAttrName, + DingTalkConnectSyncDeptAttrName: updatedSettings.DingTalkConnectSyncDeptAttrName, + WeChatConnectEnabled: updatedSettings.WeChatConnectEnabled, + WeChatConnectAppID: updatedSettings.WeChatConnectAppID, + WeChatConnectAppSecretConfigured: updatedSettings.WeChatConnectAppSecretConfigured, + WeChatConnectOpenAppID: updatedSettings.WeChatConnectOpenAppID, + WeChatConnectOpenAppSecretConfigured: updatedSettings.WeChatConnectOpenAppSecretConfigured, + WeChatConnectMPAppID: updatedSettings.WeChatConnectMPAppID, + WeChatConnectMPAppSecretConfigured: updatedSettings.WeChatConnectMPAppSecretConfigured, + WeChatConnectMobileAppID: updatedSettings.WeChatConnectMobileAppID, + WeChatConnectMobileAppSecretConfigured: updatedSettings.WeChatConnectMobileAppSecretConfigured, + WeChatConnectOpenEnabled: updatedSettings.WeChatConnectOpenEnabled, + WeChatConnectMPEnabled: updatedSettings.WeChatConnectMPEnabled, + WeChatConnectMobileEnabled: updatedSettings.WeChatConnectMobileEnabled, + WeChatConnectMode: updatedSettings.WeChatConnectMode, + WeChatConnectScopes: updatedSettings.WeChatConnectScopes, + WeChatConnectRedirectURL: updatedSettings.WeChatConnectRedirectURL, + WeChatConnectFrontendRedirectURL: updatedSettings.WeChatConnectFrontendRedirectURL, + OIDCConnectEnabled: updatedSettings.OIDCConnectEnabled, + OIDCConnectProviderName: updatedSettings.OIDCConnectProviderName, + OIDCConnectClientID: updatedSettings.OIDCConnectClientID, + OIDCConnectClientSecretConfigured: updatedSettings.OIDCConnectClientSecretConfigured, + OIDCConnectIssuerURL: updatedSettings.OIDCConnectIssuerURL, + OIDCConnectDiscoveryURL: updatedSettings.OIDCConnectDiscoveryURL, + OIDCConnectAuthorizeURL: updatedSettings.OIDCConnectAuthorizeURL, + OIDCConnectTokenURL: updatedSettings.OIDCConnectTokenURL, + OIDCConnectUserInfoURL: updatedSettings.OIDCConnectUserInfoURL, + OIDCConnectJWKSURL: updatedSettings.OIDCConnectJWKSURL, + OIDCConnectScopes: updatedSettings.OIDCConnectScopes, + OIDCConnectRedirectURL: updatedSettings.OIDCConnectRedirectURL, + OIDCConnectFrontendRedirectURL: updatedSettings.OIDCConnectFrontendRedirectURL, + OIDCConnectTokenAuthMethod: updatedSettings.OIDCConnectTokenAuthMethod, + OIDCConnectUsePKCE: updatedSettings.OIDCConnectUsePKCE, + OIDCConnectValidateIDToken: updatedSettings.OIDCConnectValidateIDToken, + OIDCConnectAllowedSigningAlgs: updatedSettings.OIDCConnectAllowedSigningAlgs, + OIDCConnectClockSkewSeconds: updatedSettings.OIDCConnectClockSkewSeconds, + OIDCConnectRequireEmailVerified: updatedSettings.OIDCConnectRequireEmailVerified, + OIDCConnectUserInfoEmailPath: updatedSettings.OIDCConnectUserInfoEmailPath, + OIDCConnectUserInfoIDPath: updatedSettings.OIDCConnectUserInfoIDPath, + OIDCConnectUserInfoUsernamePath: updatedSettings.OIDCConnectUserInfoUsernamePath, + GitHubOAuthEnabled: updatedSettings.GitHubOAuthEnabled, + GitHubOAuthClientID: updatedSettings.GitHubOAuthClientID, + GitHubOAuthClientSecretConfigured: updatedSettings.GitHubOAuthClientSecretConfigured, + GitHubOAuthRedirectURL: updatedSettings.GitHubOAuthRedirectURL, + GitHubOAuthFrontendRedirectURL: updatedSettings.GitHubOAuthFrontendRedirectURL, + GoogleOAuthEnabled: updatedSettings.GoogleOAuthEnabled, + GoogleOAuthClientID: updatedSettings.GoogleOAuthClientID, + GoogleOAuthClientSecretConfigured: updatedSettings.GoogleOAuthClientSecretConfigured, + GoogleOAuthRedirectURL: updatedSettings.GoogleOAuthRedirectURL, + GoogleOAuthFrontendRedirectURL: updatedSettings.GoogleOAuthFrontendRedirectURL, + SiteName: updatedSettings.SiteName, + SiteLogo: updatedSettings.SiteLogo, + SiteSubtitle: updatedSettings.SiteSubtitle, + APIBaseURL: updatedSettings.APIBaseURL, + ContactInfo: updatedSettings.ContactInfo, + DocURL: updatedSettings.DocURL, + HomeContent: updatedSettings.HomeContent, + HideCcsImportButton: updatedSettings.HideCcsImportButton, + PurchaseSubscriptionEnabled: updatedSettings.PurchaseSubscriptionEnabled, + PurchaseSubscriptionURL: updatedSettings.PurchaseSubscriptionURL, + TableDefaultPageSize: updatedSettings.TableDefaultPageSize, + TablePageSizeOptions: updatedSettings.TablePageSizeOptions, + CustomMenuItems: dto.ParseCustomMenuItems(updatedSettings.CustomMenuItems), + CustomEndpoints: dto.ParseCustomEndpoints(updatedSettings.CustomEndpoints), + DefaultConcurrency: updatedSettings.DefaultConcurrency, + DefaultBalance: updatedSettings.DefaultBalance, + AffiliateRebateRate: updatedSettings.AffiliateRebateRate, + AffiliateRebateFreezeHours: updatedSettings.AffiliateRebateFreezeHours, + AffiliateRebateDurationDays: updatedSettings.AffiliateRebateDurationDays, + AffiliateRebatePerInviteeCap: updatedSettings.AffiliateRebatePerInviteeCap, + DefaultUserRPMLimit: updatedSettings.DefaultUserRPMLimit, + DefaultSubscriptions: updatedDefaultSubscriptions, + EnableModelFallback: updatedSettings.EnableModelFallback, + FallbackModelAnthropic: updatedSettings.FallbackModelAnthropic, + FallbackModelOpenAI: updatedSettings.FallbackModelOpenAI, + FallbackModelGemini: updatedSettings.FallbackModelGemini, + FallbackModelAntigravity: updatedSettings.FallbackModelAntigravity, + EnableIdentityPatch: updatedSettings.EnableIdentityPatch, + IdentityPatchPrompt: updatedSettings.IdentityPatchPrompt, + OpsMonitoringEnabled: updatedSettings.OpsMonitoringEnabled, + OpsRealtimeMonitoringEnabled: updatedSettings.OpsRealtimeMonitoringEnabled, + OpsQueryModeDefault: updatedSettings.OpsQueryModeDefault, + OpsMetricsIntervalSeconds: updatedSettings.OpsMetricsIntervalSeconds, + MinClaudeCodeVersion: updatedSettings.MinClaudeCodeVersion, + MaxClaudeCodeVersion: updatedSettings.MaxClaudeCodeVersion, + AllowUngroupedKeyScheduling: updatedSettings.AllowUngroupedKeyScheduling, + BackendModeEnabled: updatedSettings.BackendModeEnabled, + EnableFingerprintUnification: updatedSettings.EnableFingerprintUnification, + EnableMetadataPassthrough: updatedSettings.EnableMetadataPassthrough, + EnableCCHSigning: updatedSettings.EnableCCHSigning, + EnableClaudeOAuthSystemPromptInjection: updatedSettings.EnableClaudeOAuthSystemPromptInjection, + ClaudeOAuthSystemPrompt: updatedSettings.ClaudeOAuthSystemPrompt, + ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks, + EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection, + RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl, + EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization, + AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion, + OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent, + MinCodexVersion: updatedSettings.MinCodexVersion, + MaxCodexVersion: updatedSettings.MaxCodexVersion, + CodexCLIOnlyBlacklist: updatedSettings.CodexCLIOnlyBlacklist, + CodexCLIOnlyWhitelist: updatedSettings.CodexCLIOnlyWhitelist, + CodexCLIOnlyAllowAppServerClients: updatedSettings.CodexCLIOnlyAllowAppServerClients, + CodexCLIOnlyEngineFingerprintSignals: updatedSettings.CodexCLIOnlyEngineFingerprintSignals, + PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource, + PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource, + PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled, + PaymentVisibleMethodWxpayEnabled: updatedSettings.PaymentVisibleMethodWxpayEnabled, + OpenAIAdvancedSchedulerEnabled: updatedSettings.OpenAIAdvancedSchedulerEnabled, + OpenAIAdvancedSchedulerStickyWeightedEnabled: updatedSettings.OpenAIAdvancedSchedulerStickyWeightedEnabled, + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: updatedSettings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled, + OpenAIAdvancedSchedulerLBTopK: updatedSettings.OpenAIAdvancedSchedulerLBTopK, + OpenAIAdvancedSchedulerWeightPriority: updatedSettings.OpenAIAdvancedSchedulerWeightPriority, + OpenAIAdvancedSchedulerWeightLoad: updatedSettings.OpenAIAdvancedSchedulerWeightLoad, + OpenAIAdvancedSchedulerWeightQueue: updatedSettings.OpenAIAdvancedSchedulerWeightQueue, + OpenAIAdvancedSchedulerWeightErrorRate: updatedSettings.OpenAIAdvancedSchedulerWeightErrorRate, + OpenAIAdvancedSchedulerWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerWeightTTFT, + OpenAIAdvancedSchedulerWeightReset: updatedSettings.OpenAIAdvancedSchedulerWeightReset, + OpenAIAdvancedSchedulerWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, + OpenAIAdvancedSchedulerWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerWeightPreviousResponse, + OpenAIAdvancedSchedulerWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerWeightSessionSticky, + OpenAIAdvancedSchedulerEffectiveLBTopK: updatedSettings.OpenAIAdvancedSchedulerEffectiveLBTopK, + OpenAIAdvancedSchedulerEffectiveWeightPriority: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightPriority, + OpenAIAdvancedSchedulerEffectiveWeightLoad: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightLoad, + OpenAIAdvancedSchedulerEffectiveWeightQueue: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightQueue, + OpenAIAdvancedSchedulerEffectiveWeightErrorRate: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightErrorRate, + OpenAIAdvancedSchedulerEffectiveWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightTTFT, + OpenAIAdvancedSchedulerEffectiveWeightReset: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightReset, + OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom, + OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse, + OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky, + BalanceLowNotifyEnabled: updatedSettings.BalanceLowNotifyEnabled, + BalanceLowNotifyThreshold: updatedSettings.BalanceLowNotifyThreshold, + BalanceLowNotifyRechargeURL: updatedSettings.BalanceLowNotifyRechargeURL, + SubscriptionExpiryNotifyEnabled: updatedSettings.SubscriptionExpiryNotifyEnabled, + AccountQuotaNotifyEnabled: updatedSettings.AccountQuotaNotifyEnabled, + AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(updatedSettings.AccountQuotaNotifyEmails), + PaymentEnabled: updatedPaymentCfg.Enabled, + PaymentMinAmount: updatedPaymentCfg.MinAmount, + PaymentMaxAmount: updatedPaymentCfg.MaxAmount, + PaymentDailyLimit: updatedPaymentCfg.DailyLimit, + PaymentOrderTimeoutMin: updatedPaymentCfg.OrderTimeoutMin, + PaymentMaxPendingOrders: updatedPaymentCfg.MaxPendingOrders, + PaymentEnabledTypes: updatedPaymentCfg.EnabledTypes, + PaymentBalanceDisabled: updatedPaymentCfg.BalanceDisabled, + PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier, + PaymentSubscriptionUSDToCNYRate: updatedPaymentCfg.SubscriptionUSDToCNYRate, + PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate, + PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy, + PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix, + PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix, + PaymentHelpImageURL: updatedPaymentCfg.HelpImageURL, + PaymentHelpText: updatedPaymentCfg.HelpText, + PaymentCancelRateLimitEnabled: updatedPaymentCfg.CancelRateLimitEnabled, + PaymentCancelRateLimitMax: updatedPaymentCfg.CancelRateLimitMax, + PaymentCancelRateLimitWindow: updatedPaymentCfg.CancelRateLimitWindow, + PaymentCancelRateLimitUnit: updatedPaymentCfg.CancelRateLimitUnit, + PaymentCancelRateLimitMode: updatedPaymentCfg.CancelRateLimitMode, + PaymentAlipayForceQRCode: updatedPaymentCfg.AlipayForceQRCode, ChannelMonitorEnabled: updatedSettings.ChannelMonitorEnabled, ChannelMonitorDefaultIntervalSeconds: updatedSettings.ChannelMonitorDefaultIntervalSeconds, @@ -2238,7 +2320,8 @@ func hasPaymentFields(req UpdateSettingsRequest) bool { req.PaymentMaxAmount != nil || req.PaymentDailyLimit != nil || req.PaymentOrderTimeoutMin != nil || req.PaymentMaxPendingOrders != nil || req.PaymentEnabledTypes != nil || req.PaymentBalanceDisabled != nil || - req.PaymentBalanceRechargeMultiplier != nil || req.PaymentRechargeFeeRate != nil || + req.PaymentBalanceRechargeMultiplier != nil || req.PaymentSubscriptionUSDToCNYRate != nil || + req.PaymentRechargeFeeRate != nil || req.PaymentLoadBalanceStrat != nil || req.PaymentProductNamePrefix != nil || req.PaymentProductNameSuffix != nil || req.PaymentHelpImageURL != nil || req.PaymentHelpText != nil || req.PaymentCancelRateLimitEnabled != nil || @@ -2677,6 +2760,42 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.OpenAIAdvancedSchedulerEnabled != after.OpenAIAdvancedSchedulerEnabled { changed = append(changed, "openai_advanced_scheduler_enabled") } + if before.OpenAIAdvancedSchedulerStickyWeightedEnabled != after.OpenAIAdvancedSchedulerStickyWeightedEnabled { + changed = append(changed, "openai_advanced_scheduler_sticky_weighted_enabled") + } + if before.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled != after.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled { + changed = append(changed, "openai_advanced_scheduler_subscription_priority_enabled") + } + if before.OpenAIAdvancedSchedulerLBTopK != after.OpenAIAdvancedSchedulerLBTopK { + changed = append(changed, "openai_advanced_scheduler_lb_top_k") + } + if before.OpenAIAdvancedSchedulerWeightPriority != after.OpenAIAdvancedSchedulerWeightPriority { + changed = append(changed, "openai_advanced_scheduler_weight_priority") + } + if before.OpenAIAdvancedSchedulerWeightLoad != after.OpenAIAdvancedSchedulerWeightLoad { + changed = append(changed, "openai_advanced_scheduler_weight_load") + } + if before.OpenAIAdvancedSchedulerWeightQueue != after.OpenAIAdvancedSchedulerWeightQueue { + changed = append(changed, "openai_advanced_scheduler_weight_queue") + } + if before.OpenAIAdvancedSchedulerWeightErrorRate != after.OpenAIAdvancedSchedulerWeightErrorRate { + changed = append(changed, "openai_advanced_scheduler_weight_error_rate") + } + if before.OpenAIAdvancedSchedulerWeightTTFT != after.OpenAIAdvancedSchedulerWeightTTFT { + changed = append(changed, "openai_advanced_scheduler_weight_ttft") + } + if before.OpenAIAdvancedSchedulerWeightReset != after.OpenAIAdvancedSchedulerWeightReset { + changed = append(changed, "openai_advanced_scheduler_weight_reset") + } + if before.OpenAIAdvancedSchedulerWeightQuotaHeadroom != after.OpenAIAdvancedSchedulerWeightQuotaHeadroom { + changed = append(changed, "openai_advanced_scheduler_weight_quota_headroom") + } + if before.OpenAIAdvancedSchedulerWeightPreviousResponse != after.OpenAIAdvancedSchedulerWeightPreviousResponse { + changed = append(changed, "openai_advanced_scheduler_weight_previous_response") + } + if before.OpenAIAdvancedSchedulerWeightSessionSticky != after.OpenAIAdvancedSchedulerWeightSessionSticky { + changed = append(changed, "openai_advanced_scheduler_weight_session_sticky") + } // 余额、订阅到期与账号限额通知 if before.BalanceLowNotifyEnabled != after.BalanceLowNotifyEnabled { changed = append(changed, "balance_low_notify_enabled") @@ -3829,3 +3948,10 @@ func equalPlatformQuotaSettings(before, after map[string]*service.DefaultPlatfor } return true } + +func stringSetting(value *string, fallback string) string { + if value == nil { + return fallback + } + return *value +} diff --git a/backend/internal/handler/admin/setting_handler_auth_source_defaults_test.go b/backend/internal/handler/admin/setting_handler_auth_source_defaults_test.go index f953f76760..1626007f19 100644 --- a/backend/internal/handler/admin/setting_handler_auth_source_defaults_test.go +++ b/backend/internal/handler/admin/setting_handler_auth_source_defaults_test.go @@ -217,12 +217,13 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS handler := NewSettingHandler(svc, nil, nil, nil, nil, nil, nil) body := map[string]any{ - "promo_code_enabled": true, - "payment_visible_method_alipay_source": "easypay", - "payment_visible_method_wxpay_source": "wxpay", - "payment_visible_method_alipay_enabled": true, - "payment_visible_method_wxpay_enabled": false, - "openai_advanced_scheduler_enabled": true, + "promo_code_enabled": true, + "payment_visible_method_alipay_source": "easypay", + "payment_visible_method_wxpay_source": "wxpay", + "payment_visible_method_alipay_enabled": true, + "payment_visible_method_wxpay_enabled": false, + "openai_advanced_scheduler_enabled": true, + "openai_advanced_scheduler_subscription_priority_enabled": true, } rawBody, err := json.Marshal(body) require.NoError(t, err) @@ -240,6 +241,7 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS require.Equal(t, "true", repo.values[service.SettingPaymentVisibleMethodAlipayEnabled]) require.Equal(t, "false", repo.values[service.SettingPaymentVisibleMethodWxpayEnabled]) require.Equal(t, "true", repo.values["openai_advanced_scheduler_enabled"]) + require.Equal(t, "true", repo.values[service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]) var resp response.Response require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) @@ -250,6 +252,7 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS require.Equal(t, true, data["payment_visible_method_alipay_enabled"]) require.Equal(t, false, data["payment_visible_method_wxpay_enabled"]) require.Equal(t, true, data["openai_advanced_scheduler_enabled"]) + require.Equal(t, true, data["openai_advanced_scheduler_subscription_priority_enabled"]) } func TestSettingHandler_UpdateSettings_PreservesLegacyBlankPaymentVisibleMethodSource(t *testing.T) { diff --git a/backend/internal/handler/dto/api_key_mapper_last_used_test.go b/backend/internal/handler/dto/api_key_mapper_last_used_test.go index 99644ced7f..d63baba91a 100644 --- a/backend/internal/handler/dto/api_key_mapper_last_used_test.go +++ b/backend/internal/handler/dto/api_key_mapper_last_used_test.go @@ -11,18 +11,20 @@ import ( func TestAPIKeyFromService_MapsLastUsedAt(t *testing.T) { lastUsed := time.Now().UTC().Truncate(time.Second) src := &service.APIKey{ - ID: 1, - UserID: 2, - Key: "sk-map-last-used", - Name: "Mapper", - Status: service.StatusActive, - LastUsedAt: &lastUsed, + ID: 1, + UserID: 2, + Key: "sk-map-last-used", + Name: "Mapper", + Status: service.StatusActive, + LastUsedAt: &lastUsed, + CurrentConcurrency: 3, } out := APIKeyFromService(src) require.NotNil(t, out) require.NotNil(t, out.LastUsedAt) require.WithinDuration(t, lastUsed, *out.LastUsedAt, time.Second) + require.Equal(t, 3, out.CurrentConcurrency) } func TestAPIKeyFromService_MapsNilLastUsedAt(t *testing.T) { diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 5bbab4d45f..2b5bdccad7 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -79,31 +79,32 @@ func APIKeyFromService(k *service.APIKey) *APIKey { return nil } out := &APIKey{ - ID: k.ID, - UserID: k.UserID, - Key: k.Key, - Name: k.Name, - GroupID: k.GroupID, - Status: k.Status, - IPWhitelist: k.IPWhitelist, - IPBlacklist: k.IPBlacklist, - LastUsedAt: k.LastUsedAt, - Quota: k.Quota, - QuotaUsed: k.QuotaUsed, - ExpiresAt: k.ExpiresAt, - CreatedAt: k.CreatedAt, - UpdatedAt: k.UpdatedAt, - RateLimit5h: k.RateLimit5h, - RateLimit1d: k.RateLimit1d, - RateLimit7d: k.RateLimit7d, - Usage5h: k.EffectiveUsage5h(), - Usage1d: k.EffectiveUsage1d(), - Usage7d: k.EffectiveUsage7d(), - Window5hStart: k.Window5hStart, - Window1dStart: k.Window1dStart, - Window7dStart: k.Window7dStart, - User: UserFromServiceShallow(k.User), - Group: GroupFromServiceShallow(k.Group), + ID: k.ID, + UserID: k.UserID, + Key: k.Key, + Name: k.Name, + GroupID: k.GroupID, + Status: k.Status, + IPWhitelist: k.IPWhitelist, + IPBlacklist: k.IPBlacklist, + LastUsedAt: k.LastUsedAt, + Quota: k.Quota, + QuotaUsed: k.QuotaUsed, + ExpiresAt: k.ExpiresAt, + CreatedAt: k.CreatedAt, + UpdatedAt: k.UpdatedAt, + CurrentConcurrency: k.CurrentConcurrency, + RateLimit5h: k.RateLimit5h, + RateLimit1d: k.RateLimit1d, + RateLimit7d: k.RateLimit7d, + Usage5h: k.EffectiveUsage5h(), + Usage1d: k.EffectiveUsage1d(), + Usage7d: k.EffectiveUsage7d(), + Window5hStart: k.Window5hStart, + Window1dStart: k.Window1dStart, + Window7dStart: k.Window7dStart, + User: UserFromServiceShallow(k.User), + Group: GroupFromServiceShallow(k.Group), } if k.Window5hStart != nil && !service.IsWindowExpired(k.Window5hStart, service.RateLimitWindow5h) { t := k.Window5hStart.Add(service.RateLimitWindow5h) diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 1d6f73bb99..99fba54980 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -208,7 +208,29 @@ type SystemSettings struct { PaymentVisibleMethodWxpayEnabled bool `json:"payment_visible_method_wxpay_enabled"` // OpenAI account scheduling - OpenAIAdvancedSchedulerEnabled bool `json:"openai_advanced_scheduler_enabled"` + OpenAIAdvancedSchedulerEnabled bool `json:"openai_advanced_scheduler_enabled"` + OpenAIAdvancedSchedulerStickyWeightedEnabled bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"` + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled bool `json:"openai_advanced_scheduler_subscription_priority_enabled"` + OpenAIAdvancedSchedulerLBTopK string `json:"openai_advanced_scheduler_lb_top_k"` + OpenAIAdvancedSchedulerWeightPriority string `json:"openai_advanced_scheduler_weight_priority"` + OpenAIAdvancedSchedulerWeightLoad string `json:"openai_advanced_scheduler_weight_load"` + OpenAIAdvancedSchedulerWeightQueue string `json:"openai_advanced_scheduler_weight_queue"` + OpenAIAdvancedSchedulerWeightErrorRate string `json:"openai_advanced_scheduler_weight_error_rate"` + OpenAIAdvancedSchedulerWeightTTFT string `json:"openai_advanced_scheduler_weight_ttft"` + OpenAIAdvancedSchedulerWeightReset string `json:"openai_advanced_scheduler_weight_reset"` + OpenAIAdvancedSchedulerWeightQuotaHeadroom string `json:"openai_advanced_scheduler_weight_quota_headroom"` + OpenAIAdvancedSchedulerWeightPreviousResponse string `json:"openai_advanced_scheduler_weight_previous_response"` + OpenAIAdvancedSchedulerWeightSessionSticky string `json:"openai_advanced_scheduler_weight_session_sticky"` + OpenAIAdvancedSchedulerEffectiveLBTopK string `json:"openai_advanced_scheduler_effective_lb_top_k"` + OpenAIAdvancedSchedulerEffectiveWeightPriority string `json:"openai_advanced_scheduler_effective_weight_priority"` + OpenAIAdvancedSchedulerEffectiveWeightLoad string `json:"openai_advanced_scheduler_effective_weight_load"` + OpenAIAdvancedSchedulerEffectiveWeightQueue string `json:"openai_advanced_scheduler_effective_weight_queue"` + OpenAIAdvancedSchedulerEffectiveWeightErrorRate string `json:"openai_advanced_scheduler_effective_weight_error_rate"` + OpenAIAdvancedSchedulerEffectiveWeightTTFT string `json:"openai_advanced_scheduler_effective_weight_ttft"` + OpenAIAdvancedSchedulerEffectiveWeightReset string `json:"openai_advanced_scheduler_effective_weight_reset"` + OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom string `json:"openai_advanced_scheduler_effective_weight_quota_headroom"` + OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse string `json:"openai_advanced_scheduler_effective_weight_previous_response"` + OpenAIAdvancedSchedulerEffectiveWeightSessionSticky string `json:"openai_advanced_scheduler_effective_weight_session_sticky"` // Payment configuration PaymentEnabled bool `json:"payment_enabled"` @@ -220,6 +242,7 @@ type SystemSettings struct { PaymentEnabledTypes []string `json:"payment_enabled_types"` PaymentBalanceDisabled bool `json:"payment_balance_disabled"` PaymentBalanceRechargeMultiplier float64 `json:"payment_balance_recharge_multiplier"` + PaymentSubscriptionUSDToCNYRate float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate float64 `json:"payment_recharge_fee_rate"` PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` PaymentProductNamePrefix string `json:"payment_product_name_prefix"` diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index b08dea5680..8f0d23be65 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -63,6 +63,8 @@ type APIKey struct { ExpiresAt *time.Time `json:"expires_at"` // Expiration time (nil = never expires) CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` + // CurrentConcurrency is the real-time active request count for this API key. + CurrentConcurrency int `json:"current_concurrency"` // Rate limit fields RateLimit5h float64 `json:"rate_limit_5h"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index b65dedf9cc..b20d9ef652 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -1006,7 +1006,8 @@ func (h *GatewayHandler) Models(c *gin.Context) { // Get available models from account configurations for the selected group platform. availableModels := h.gatewayService.GetAvailableModels(c.Request.Context(), groupID, platform) if apiKey != nil && apiKey.Group != nil && apiKey.Group.CustomModelsListEnabled() { - availableModels = filterModelsByCustomList(availableModels, defaultModelIDsForPlatform(platform), apiKey.Group.ModelsListConfig.Models) + fallbackModels := defaultModelIDsForPlatform(platform) + availableModels = filterModelsByCustomList(customModelsListSource(platform, availableModels, fallbackModels), fallbackModels, apiKey.Group.ModelsListConfig.Models) writeCustomModelsList(c, platform, availableModels) return } @@ -1090,6 +1091,13 @@ func writeOpenAIModelsList(c *gin.Context, modelIDs []string) { }) } +func customModelsListSource(platform string, availableModels, fallbackModels []string) []string { + if platform == service.PlatformAnthropic && len(availableModels) > 0 { + return mergeModelIDs(availableModels, fallbackModels) + } + return availableModels +} + func filterModelsByCustomList(availableModels, fallbackModels, selectedModels []string) []string { if len(selectedModels) == 0 { return availableModels @@ -1158,6 +1166,15 @@ func defaultModelIDsForPlatform(platform string) []string { ids = append(ids, model.ID) } return ids + case service.PlatformAnthropic: + ids := make([]string, 0, len(claude.DefaultModels)+len(antigravity.DefaultModels())) + for _, model := range claude.DefaultModels { + ids = append(ids, model.ID) + } + for _, model := range antigravity.DefaultModels() { + ids = append(ids, model.ID) + } + return mergeModelIDs(ids, nil) case service.PlatformGrok: return xai.DefaultModelIDs() default: @@ -1169,6 +1186,25 @@ func defaultModelIDsForPlatform(platform string) []string { } } +func mergeModelIDs(primary, secondary []string) []string { + seen := make(map[string]struct{}, len(primary)+len(secondary)) + merged := make([]string, 0, len(primary)+len(secondary)) + for _, models := range [][]string{primary, secondary} { + for _, model := range models { + model = strings.TrimSpace(model) + if model == "" { + continue + } + if _, ok := seen[model]; ok { + continue + } + seen[model] = struct{}{} + merged = append(merged, model) + } + } + return merged +} + // AntigravityModels 返回 Antigravity 支持的全部模型 // GET /antigravity/models func (h *GatewayHandler) AntigravityModels(c *gin.Context) { diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index b948ac8fc7..48110da93f 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -10,6 +10,7 @@ import ( "sync" "time" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" @@ -211,6 +212,14 @@ func (h *ConcurrencyHelper) TryAcquireUserSlot(ctx context.Context, userID int64 return result.ReleaseFunc, true, nil } +func (h *ConcurrencyHelper) TryAcquireUserSlotForAPIKey(ctx context.Context, userID int64, maxConcurrency int, apiKeyID int64) (func(), bool, error) { + releaseFunc, acquired, err := h.TryAcquireUserSlot(ctx, userID, maxConcurrency) + if err != nil || !acquired { + return releaseFunc, acquired, err + } + return h.withAPIKeySlot(ctx, apiKeyID, releaseFunc), true, nil +} + // TryAcquireAccountSlot 尝试立即获取账号并发槽位。 // 返回值: (releaseFunc, acquired, error) func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (func(), bool, error) { @@ -241,7 +250,7 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI } if acquired { - return releaseFunc, nil + return h.withAPIKeySlotFromGin(c, releaseFunc), nil } queueLimit := service.CalculateMaxWait(maxConcurrency) - maxConcurrency @@ -258,7 +267,37 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI defer h.DecrementWaitCount(ctx, userID) // Need to wait - handle streaming ping if needed - return h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false) + releaseFunc, err = h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false) + if err != nil { + return nil, err + } + return h.withAPIKeySlotFromGin(c, releaseFunc), nil +} + +func (h *ConcurrencyHelper) withAPIKeySlotFromGin(c *gin.Context, releaseFunc func()) func() { + if c == nil { + return releaseFunc + } + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil { + return releaseFunc + } + return h.withAPIKeySlot(c.Request.Context(), apiKey.ID, releaseFunc) +} + +func (h *ConcurrencyHelper) withAPIKeySlot(ctx context.Context, apiKeyID int64, releaseFunc func()) func() { + if h == nil || h.concurrencyService == nil || apiKeyID <= 0 { + return releaseFunc + } + apiKeyReleaseFunc := h.concurrencyService.TrackAPIKeySlot(ctx, apiKeyID) + return func() { + if releaseFunc != nil { + releaseFunc() + } + if apiKeyReleaseFunc != nil { + apiKeyReleaseFunc() + } + } } // AcquireAccountSlotWithWait acquires an account concurrency slot, waiting if necessary. diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go index fb17481f1c..5e0697f083 100644 --- a/backend/internal/handler/gateway_helper_hotpath_test.go +++ b/backend/internal/handler/gateway_helper_hotpath_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -29,6 +30,9 @@ type helperConcurrencyCacheStub struct { waitDecrementCalls int waitMaxWait int waitIncrementHook func() + apiKeyTrackCalls int + apiKeyReleaseCalls int + apiKeyTrackIDs []int64 } func (s *helperConcurrencyCacheStub) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { @@ -97,6 +101,29 @@ func (s *helperConcurrencyCacheStub) GetUserConcurrency(ctx context.Context, use return 0, nil } +func (s *helperConcurrencyCacheStub) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.apiKeyTrackCalls++ + s.apiKeyTrackIDs = append(s.apiKeyTrackIDs, apiKeyID) + return nil +} + +func (s *helperConcurrencyCacheStub) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.apiKeyReleaseCalls++ + return nil +} + +func (s *helperConcurrencyCacheStub) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + out := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + out[apiKeyID] = 0 + } + return out, nil +} + func (s *helperConcurrencyCacheStub) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) { s.mu.Lock() s.waitIncrementCalls++ @@ -270,6 +297,48 @@ func TestAcquireUserSlotWithWait_ImmediateAcquireSkipsWaitQueue(t *testing.T) { require.Equal(t, 1, cache.userReleaseCalls) } +func TestAcquireUserSlotWithWait_TracksAPIKeySlot(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{true}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 77}) + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted) + require.NoError(t, err) + require.NotNil(t, release) + require.Equal(t, 1, cache.apiKeyTrackCalls) + require.Equal(t, []int64{77}, cache.apiKeyTrackIDs) + + release() + + require.Equal(t, 1, cache.userReleaseCalls) + require.Equal(t, 1, cache.apiKeyReleaseCalls) +} + +func TestTryAcquireUserSlotForAPIKey_TracksAPIKeySlot(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{true}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + + release, acquired, err := helper.TryAcquireUserSlotForAPIKey(context.Background(), 202, 3, 77) + require.NoError(t, err) + require.True(t, acquired) + require.NotNil(t, release) + require.Equal(t, 1, cache.apiKeyTrackCalls) + require.Equal(t, []int64{77}, cache.apiKeyTrackIDs) + + release() + + require.Equal(t, 1, cache.userReleaseCalls) + require.Equal(t, 1, cache.apiKeyReleaseCalls) +} + func TestAcquireUserSlotWithWait_WaitSuccessDecrementsBeforeReturn(t *testing.T) { cache := &helperConcurrencyCacheStub{ userSeq: []bool{false, true}, diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index c5238f2a2f..6011e13027 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -269,6 +269,149 @@ func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMappin require.Equal(t, []string{"claude-sonnet-4-6"}, modelIDsForTest(got.Data)) } +func TestGatewayModels_AnthropicCustomModelsListIncludesOAuthClaudeAndMappedDeepSeek(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(28) + h := newGatewayModelsHandlerForTest( + &gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + }, + { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "deepseek-v4-pro": "deepseek-v4-pro", + }, + }, + }, + }, + }, + }, + ) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformAnthropic, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"claude-fable-5", "claude-opus-4-8", "deepseek-v4-pro"}, + }, + }, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"claude-fable-5", "claude-opus-4-8", "deepseek-v4-pro"}, modelIDsForTest(got.Data)) +} + +func TestGatewayModels_AnthropicCustomModelsListDisabledKeepsMappedModelList(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(29) + h := newGatewayModelsHandlerForTest( + &gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + }, + { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "deepseek-v4-pro": "deepseek-v4-pro", + }, + }, + }, + }, + }, + }, + ) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformAnthropic, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: false, + Models: []string{"claude-fable-5", "deepseek-v4-pro"}, + }, + }, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"deepseek-v4-pro"}, modelIDsForTest(got.Data)) +} + +func TestGatewayModels_AnthropicCustomModelsListIncludesOAuthClaudeWithoutMappings(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(30) + h := newGatewayModelsHandlerForTest( + &gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + }, + }, + }, + }, + ) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformAnthropic, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"claude-opus-4-6-thinking", "claude-sonnet-4-5"}, + }, + }, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"claude-opus-4-6-thinking", "claude-sonnet-4-5"}, modelIDsForTest(got.Data)) +} + func TestGatewayModels_CustomModelsListCanReturnEmptyWhenSelectionsUnavailable(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 8e236ea49f..4fd1411b23 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -174,6 +174,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. service.OpenAIUpstreamTransportHTTPSSE, "", false, + false, service.PlatformGrok, ) if err != nil { diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index ca43a2cff3..baff1dcbd6 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index a80c7f7d96..8be533c723 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -117,6 +117,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { service.OpenAIUpstreamTransportHTTPSSE, service.OpenAIEndpointCapabilityEmbeddings, false, + false, ) if err != nil { reqLog.Warn("openai_embeddings.account_select_failed", diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index ec530e8ab6..fc9c4d5df7 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -110,6 +110,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, openAICompatibleRequestPlatform(apiKey), ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 8eec667fa6..7f097afa4b 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -350,6 +350,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, requireCompact, + false, requestPlatform, ) if err != nil { @@ -783,6 +784,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { @@ -1266,6 +1268,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "previous_response_id must be a response.id (resp_*), not a message id") return } + firstMessageToolCoverage := service.AnalyzeToolCallOutputContextCoverageBytes(firstMessage) + previousResponseCanMove := !firstMessageToolCoverage.HasFunctionCallOutput || firstMessageToolCoverage.ContextCoversAllCallIDs reqLog = reqLog.With( zap.Bool("ws_ingress", true), zap.String("model", reqModel), @@ -1318,7 +1322,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // 必须尽早注册,确保任何 early return 都能释放已获取的并发槽位。 defer releaseTurnSlots() - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { reqLog.Warn("openai.websocket_user_slot_acquire_failed", zap.Error(err)) closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot") @@ -1333,7 +1337,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { if currentUserRelease != nil { return true } - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { reqLog.Warn("openai.websocket_user_slot_reacquire_failed", zap.Error(err)) closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot") @@ -1381,6 +1385,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { requiredTransport, service.OpenAIEndpointCapabilityChatCompletions, false, + previousResponseCanMove, requestPlatform, ) if err != nil { @@ -1484,7 +1489,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // 防御式清理:避免异常路径下旧槽位覆盖导致泄漏。 releaseTurnSlots() // 非首轮 turn 需要重新抢占并发槽位,避免长连接空闲占槽。 - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire user concurrency slot", err) } @@ -1581,8 +1586,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // 说明该会话链不属于本次调度到的账号,原样转发会触发上游会话链鉴权失败(“鉴权失败,请检查 API Key”)。 // 故剥离首包里的 previous_response_id,改用首包内 input 重建上下文;带 function_call_output 的 // 工具续链无法重建,保持原样。仅作用于首轮首包,后续 turn 的续链由 WS 转发层既有逻辑处理。 - if previousResponseID != "" && !scheduleDecision.StickyPreviousHit && - !service.ValidateFunctionCallOutputContextBytes(wsFirstMessage).HasFunctionCallOutput { + if previousResponseID != "" && !scheduleDecision.StickyPreviousHit && previousResponseCanMove { wsFirstMessage = service.RemovePreviousResponseIDFromBody(wsFirstMessage) reqLog.Debug("openai.websocket_previous_response_id_stripped_cross_group", zap.Int64("account_id", account.ID), diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index 7cdf73cd3d..a267d73724 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -150,6 +150,7 @@ func (h *PaymentHandler) GetCheckoutInfo(c *gin.Context) { Plans: planList, BalanceDisabled: cfg.BalanceDisabled, BalanceRechargeMultiplier: cfg.BalanceRechargeMultiplier, + SubscriptionUSDToCNYRate: cfg.SubscriptionUSDToCNYRate, RechargeFeeRate: cfg.RechargeFeeRate, HelpText: cfg.HelpText, HelpImageURL: cfg.HelpImageURL, @@ -165,6 +166,7 @@ type checkoutInfoResponse struct { Plans []checkoutPlan `json:"plans"` BalanceDisabled bool `json:"balance_disabled"` BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` + SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` HelpText string `json:"help_text"` HelpImageURL string `json:"help_image_url"` diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go index 9d0f1d8fac..be6dc917bb 100644 --- a/backend/internal/handler/usage_handler.go +++ b/backend/internal/handler/usage_handler.go @@ -322,6 +322,9 @@ func (h *UsageHandler) ListErrors(c *gin.Context) { filter.ErrorTypesAny = types } + // 排序对齐用量明细:列白名单与方向归一在 repo 层,非法值回退 created_at DESC。 + filter.SetSort(c.Query("sort_by"), c.Query("sort_order")) + result, err := h.opsService.ListUserErrorRequests(c.Request.Context(), subject.UserID, filter) if err != nil { response.ErrorFrom(c, err) diff --git a/backend/internal/payment/provider/easypay.go b/backend/internal/payment/provider/easypay.go index 32d6b7bebf..f1c17427ad 100644 --- a/backend/internal/payment/provider/easypay.go +++ b/backend/internal/payment/provider/easypay.go @@ -39,6 +39,12 @@ type EasyPay struct { httpClient *http.Client } +type easyPayCustomMethod struct { + Type string `json:"type"` + UpstreamType string `json:"upstreamType"` + DisplayName string `json:"displayName"` +} + // NewEasyPay creates a new EasyPay provider. // config keys: pid, pkey, apiBase, notifyUrl, returnUrl, cid, cidAlipay, cidWxpay func NewEasyPay(instanceID string, config map[string]string) (*EasyPay, error) { @@ -95,7 +101,13 @@ func (e *EasyPay) apiBase() string { func (e *EasyPay) Name() string { return "EasyPay" } func (e *EasyPay) ProviderKey() string { return payment.TypeEasyPay } func (e *EasyPay) SupportedTypes() []payment.PaymentType { - return []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay} + types := []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay} + for _, method := range e.customMethods() { + if method.Type != "" { + types = append(types, method.Type) + } + } + return types } func (e *EasyPay) MerchantIdentityMetadata() map[string]string { @@ -124,13 +136,14 @@ func (e *EasyPay) CreatePayment(ctx context.Context, req payment.CreatePaymentRe // TradeNo is empty; it arrives via the notify callback after payment. func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) { notifyURL, returnURL := e.resolveURLs(req) + paymentType := e.upstreamPaymentType(req.PaymentType) params := map[string]string{ - "pid": e.config["pid"], "type": req.PaymentType, + "pid": e.config["pid"], "type": paymentType, "out_trade_no": req.OrderID, "notify_url": notifyURL, "return_url": returnURL, "name": req.Subject, "money": req.Amount, } - if cid := e.resolveCID(req.PaymentType); cid != "" { + if cid := e.resolveCID(paymentType); cid != "" { params["cid"] = cid } if req.IsMobile { @@ -150,13 +163,14 @@ func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*paym // createAPIPayment calls mapi.php to get payurl/qrcode (existing behavior). func (e *EasyPay) createAPIPayment(ctx context.Context, req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) { notifyURL, returnURL := e.resolveURLs(req) + paymentType := e.upstreamPaymentType(req.PaymentType) params := map[string]string{ - "pid": e.config["pid"], "type": req.PaymentType, + "pid": e.config["pid"], "type": paymentType, "out_trade_no": req.OrderID, "notify_url": notifyURL, "return_url": returnURL, "name": req.Subject, "money": req.Amount, "clientip": req.ClientIP, } - if cid := e.resolveCID(req.PaymentType); cid != "" { + if cid := e.resolveCID(paymentType); cid != "" { params["cid"] = cid } if req.IsMobile { @@ -204,6 +218,41 @@ func (e *EasyPay) resolveURLs(req payment.CreatePaymentRequest) (string, string) return notifyURL, returnURL } +func (e *EasyPay) customMethods() []easyPayCustomMethod { + if e == nil { + return nil + } + raw := strings.TrimSpace(e.config["customMethods"]) + if raw == "" { + return nil + } + var methods []easyPayCustomMethod + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return nil + } + result := make([]easyPayCustomMethod, 0, len(methods)) + for _, method := range methods { + method.Type = strings.TrimSpace(method.Type) + method.UpstreamType = strings.TrimSpace(method.UpstreamType) + method.DisplayName = strings.TrimSpace(method.DisplayName) + if method.Type == "" || method.UpstreamType == "" { + continue + } + result = append(result, method) + } + return result +} + +func (e *EasyPay) upstreamPaymentType(paymentType string) string { + paymentType = strings.TrimSpace(paymentType) + for _, method := range e.customMethods() { + if paymentType == method.Type { + return method.UpstreamType + } + } + return paymentType +} + func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.QueryOrderResponse, error) { params := map[string]string{ "act": "order", "pid": e.config["pid"], diff --git a/backend/internal/payment/provider/easypay_refund_test.go b/backend/internal/payment/provider/easypay_refund_test.go index 9e0e4942c2..3b76329870 100644 --- a/backend/internal/payment/provider/easypay_refund_test.go +++ b/backend/internal/payment/provider/easypay_refund_test.go @@ -179,6 +179,102 @@ func TestEasyPayRefundResponseErrors(t *testing.T) { } } +func TestEasyPayCustomMethodsUseConfiguredUpstreamType(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "paymentMode": paymentModePopup, + "customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{ + OrderID: "sub2-custom-1", + Amount: "1.00", + PaymentType: "usdt_trc20", + Subject: "Custom EasyPay", + }) + if err != nil { + t.Fatalf("CreatePayment: %v", err) + } + payURL, err := url.Parse(resp.PayURL) + if err != nil { + t.Fatalf("parse pay url: %v", err) + } + if got := payURL.Query().Get("type"); got != "usdt" { + t.Fatalf("pay url type = %q, want usdt (%s)", got, resp.PayURL) + } +} + +func TestEasyPayCustomMethodsResolveCIDFromConfiguredUpstreamType(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "paymentMode": paymentModePopup, + "cidAlipay": "cid-alipay", + "cidWxpay": "cid-wxpay", + "customMethods": `[{"type":"ldc","upstreamType":"alipay","displayName":"LDC"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{ + OrderID: "sub2-custom-cid", + Amount: "1.00", + PaymentType: "ldc", + Subject: "Custom EasyPay CID", + }) + if err != nil { + t.Fatalf("CreatePayment: %v", err) + } + payURL, err := url.Parse(resp.PayURL) + if err != nil { + t.Fatalf("parse pay url: %v", err) + } + if got := payURL.Query().Get("type"); got != "alipay" { + t.Fatalf("pay url type = %q, want alipay (%s)", got, resp.PayURL) + } + if got := payURL.Query().Get("cid"); got != "cid-alipay" { + t.Fatalf("pay url cid = %q, want cid-alipay (%s)", got, resp.PayURL) + } +} + +func TestEasyPaySupportedTypesIncludeCustomMethods(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + got := strings.Join(provider.SupportedTypes(), ",") + for _, want := range []string{"alipay", "wxpay", "ldc", "usdt_trc20"} { + if !strings.Contains(got, want) { + t.Fatalf("SupportedTypes() = %q, want it to include %q", got, want) + } + } +} + func newTestEasyPay(t *testing.T, apiBase string) *EasyPay { t.Helper() diff --git a/backend/internal/pkg/openai/constants.go b/backend/internal/pkg/openai/constants.go index f658cf0675..c9d391df4e 100644 --- a/backend/internal/pkg/openai/constants.go +++ b/backend/internal/pkg/openai/constants.go @@ -18,6 +18,9 @@ type Model struct { // DefaultModels OpenAI models list var DefaultModels = []Model{ + {ID: "gpt-5.6-sol", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Sol"}, + {ID: "gpt-5.6-terra", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Terra"}, + {ID: "gpt-5.6-luna", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Luna"}, {ID: "gpt-5.5", Object: "model", Created: 1776873600, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.5"}, {ID: "gpt-5.4", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4"}, {ID: "gpt-5.4-mini", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4 Mini"}, diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 19a1ff0485..a8015ddba0 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -482,7 +482,7 @@ func (r *accountRepository) List(ctx context.Context, params pagination.Paginati return r.ListWithFilters(ctx, params, "", "", "", "", 0, "") } -func (r *accountRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) { +func (r *accountRepository) accountListFilteredQuery(platform, accountType, status, search string, groupID int64, privacyMode string) *dbent.AccountQuery { q := r.client.Account.Query() if platform != "" { @@ -575,6 +575,11 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati })) } + return q +} + +func (r *accountRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) { + q := r.accountListFilteredQuery(platform, accountType, status, search, groupID, privacyMode) // Clone before Count so interceptor-appended predicates (SoftDeleteMixin's // deleted_at IS NULL) don't accumulate on the shared builder and pollute the // subsequent list query. Same pattern used in group_repo/promo_code_repo/user_repo @@ -603,6 +608,14 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati return outAccounts, paginationResultFromTotal(int64(total), params), nil } +func (r *accountRepository) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) { + accounts, err := r.accountListFilteredQuery(platform, accountType, status, search, groupID, privacyMode).All(ctx) + if err != nil { + return nil, err + } + return r.accountsToService(ctx, accounts) +} + func (r *accountRepository) ListOpsAccountsForStats(ctx context.Context, platformFilter string, groupIDFilter *int64) ([]service.Account, error) { if r == nil || r.client == nil { return []service.Account{}, nil @@ -1061,6 +1074,90 @@ func (r *accountRepository) ListSchedulableByGroupID(ctx context.Context, groupI }) } +func (r *accountRepository) ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]service.GroupAccountCapacityRow, error) { + groupIDs = uniquePositiveInt64s(groupIDs) + if len(groupIDs) == 0 { + return []service.GroupAccountCapacityRow{}, nil + } + if r.sql == nil { + rows := make([]service.GroupAccountCapacityRow, 0) + for _, groupID := range groupIDs { + accounts, err := r.ListSchedulableByGroupID(ctx, groupID) + if err != nil { + return nil, err + } + for i := range accounts { + acc := &accounts[i] + rows = append(rows, service.GroupAccountCapacityRow{ + GroupID: groupID, + AccountID: acc.ID, + Concurrency: acc.Concurrency, + Extra: copyJSONMap(acc.Extra), + SessionWindowStart: acc.SessionWindowStart, + SessionWindowEnd: acc.SessionWindowEnd, + SessionWindowStatus: acc.SessionWindowStatus, + }) + } + } + return rows, nil + } + + rows, err := r.sql.QueryContext(ctx, ` + SELECT + ag.group_id, + a.id AS account_id, + a.concurrency, + COALESCE(a.extra, '{}'::jsonb)::text AS extra, + a.session_window_start, + a.session_window_end, + COALESCE(a.session_window_status, '') AS session_window_status + FROM account_groups ag + JOIN accounts a ON a.id = ag.account_id + WHERE ag.group_id = ANY($1) + AND a.deleted_at IS NULL + AND a.status = $2 + AND a.schedulable = TRUE + AND (a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= $3) + AND (a.expires_at IS NULL OR a.expires_at > $3 OR a.auto_pause_on_expired = FALSE) + AND (a.overload_until IS NULL OR a.overload_until <= $3) + AND (a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= $3) + ORDER BY ag.group_id ASC, ag.priority ASC, a.priority ASC, a.id ASC + `, pq.Array(groupIDs), service.StatusActive, time.Now()) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + out := make([]service.GroupAccountCapacityRow, 0) + for rows.Next() { + var row service.GroupAccountCapacityRow + var extraRaw string + if err := rows.Scan( + &row.GroupID, + &row.AccountID, + &row.Concurrency, + &extraRaw, + &row.SessionWindowStart, + &row.SessionWindowEnd, + &row.SessionWindowStatus, + ); err != nil { + return nil, err + } + if extraRaw != "" && extraRaw != "null" { + var extra map[string]any + if err := json.Unmarshal([]byte(extraRaw), &extra); err != nil { + return nil, err + } + row.Extra = extra + } + out = append(out, row) + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} + func (r *accountRepository) ListSchedulableByPlatform(ctx context.Context, platform string) ([]service.Account, error) { now := time.Now() accounts, err := r.client.Account.Query(). diff --git a/backend/internal/repository/concurrency_cache.go b/backend/internal/repository/concurrency_cache.go index 8cdb5015de..9e8e216b9e 100644 --- a/backend/internal/repository/concurrency_cache.go +++ b/backend/internal/repository/concurrency_cache.go @@ -27,6 +27,8 @@ const ( accountSlotKeyPrefix = "concurrency:account:" // 格式: concurrency:user:{userID} userSlotKeyPrefix = "concurrency:user:" + // 格式: concurrency:api_key:{apiKeyID} + apiKeySlotKeyPrefix = "concurrency:api_key:" // 等待队列计数器格式: concurrency:wait:{userID} waitQueueKeyPrefix = "concurrency:wait:" // 账号级等待队列计数器格式: wait:account:{accountID} @@ -108,6 +110,28 @@ var ( return redis.call('ZCARD', key) `) + // trackSlotScript 记录 stats-only 槽位,不做并发上限判断。 + // KEYS[1] = 有序集合键 + // ARGV[1] = TTL(秒) + // ARGV[2] = requestID + trackSlotScript = redis.NewScript(` + -- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME') + -- replicates correctly. No-op on Redis 5.0+ (effects replication is default). + redis.replicate_commands() + local key = KEYS[1] + local ttl = tonumber(ARGV[1]) + local requestID = ARGV[2] + + local timeResult = redis.call('TIME') + local now = tonumber(timeResult[1]) + local expireBefore = now - ttl + + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + redis.call('ZADD', key, now, requestID) + redis.call('EXPIRE', key, ttl) + return 1 + `) + // incrementWaitScript - refreshes TTL on each increment to keep queue depth accurate // KEYS[1] = wait queue key // ARGV[1] = maxWait @@ -237,6 +261,10 @@ func userSlotKey(userID int64) string { return fmt.Sprintf("%s%d", userSlotKeyPrefix, userID) } +func apiKeySlotKey(apiKeyID int64) string { + return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) +} + func waitQueueKey(userID int64) string { return fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) } @@ -546,6 +574,54 @@ func (c *concurrencyCache) GetUserConcurrency(ctx context.Context, userID int64) return result, nil } +func (c *concurrencyCache) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + key := apiKeySlotKey(apiKeyID) + _, err := trackSlotScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds, requestID).Result() + return err +} + +func (c *concurrencyCache) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + key := apiKeySlotKey(apiKeyID) + return c.rdb.ZRem(ctx, key, requestID).Err() +} + +func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + if len(apiKeyIDs) == 0 { + return map[int64]int{}, nil + } + + now, err := c.rdb.Time(ctx).Result() + if err != nil { + return nil, fmt.Errorf("redis TIME: %w", err) + } + cutoffTime := now.Unix() - int64(c.slotTTLSeconds) + + pipe := c.rdb.Pipeline() + type apiKeyCmd struct { + apiKeyID int64 + zcardCmd *redis.IntCmd + } + cmds := make([]apiKeyCmd, 0, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + slotKey := apiKeySlotKeyPrefix + strconv.FormatInt(apiKeyID, 10) + pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + cmds = append(cmds, apiKeyCmd{ + apiKeyID: apiKeyID, + zcardCmd: pipe.ZCard(ctx, slotKey), + }) + } + + if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) { + return nil, fmt.Errorf("pipeline exec: %w", err) + } + + result := make(map[int64]int, len(apiKeyIDs)) + for _, cmd := range cmds { + result[cmd.apiKeyID] = int(cmd.zcardCmd.Val()) + } + return result, nil +} + // Wait queue operations func (c *concurrencyCache) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) { @@ -814,6 +890,8 @@ func (c *concurrencyCache) CleanupExpiredAccountSlotKeys(ctx context.Context) er // CleanupStaleProcessSlots 启动时清理非当前进程前缀的槽位。 // 清理范围来自活跃索引,避免在 Redis 上 SCAN 全部 concurrency:* 键。 +// API Key 槽位(concurrency:api_key:*)是 stats-only 数据:每次 Track/读取都会按分数 +// 裁剪过期成员,key 自带 TTL,可在一个 slot TTL 内自愈,因此不参与启动清理。 func (c *concurrencyCache) CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error { if activeRequestPrefix == "" { return nil diff --git a/backend/internal/repository/concurrency_cache_integration_test.go b/backend/internal/repository/concurrency_cache_integration_test.go index aa18ba16cb..3c831487de 100644 --- a/backend/internal/repository/concurrency_cache_integration_test.go +++ b/backend/internal/repository/concurrency_cache_integration_test.go @@ -3,6 +3,7 @@ package repository import ( + "context" "errors" "fmt" "strconv" @@ -37,6 +38,18 @@ func (s *ConcurrencyCacheSuite) SetupTest() { s.cache = s.rawCache } +type apiKeyConcurrencyCacheForTest interface { + TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) +} + +func (s *ConcurrencyCacheSuite) apiKeyConcurrencyCache() apiKeyConcurrencyCacheForTest { + cache, ok := s.cache.(apiKeyConcurrencyCacheForTest) + require.True(s.T(), ok) + return cache +} + func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() { accountID := int64(10) reqID1, reqID2, reqID3 := "req1", "req2", "req3" @@ -218,6 +231,34 @@ func (s *ConcurrencyCacheSuite) TestUserSlot_TTL() { s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL) } +func (s *ConcurrencyCacheSuite) TestAPIKeySlot_TrackReleaseAndBatchCount() { + cache := s.apiKeyConcurrencyCache() + apiKeyID := int64(300) + emptyAPIKeyID := int64(301) + slotKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) + + require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req1")) + require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req2")) + + counts, err := cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID, emptyAPIKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), map[int64]int{apiKeyID: 2, emptyAPIKeyID: 0}, counts) + + ttl, err := s.rdb.TTL(s.ctx, slotKey).Result() + require.NoError(s.T(), err, "TTL") + s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL) + + require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req1")) + counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), 1, counts[apiKeyID]) + + require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req2")) + counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), 0, counts[apiKeyID]) +} + func (s *ConcurrencyCacheSuite) TestWaitQueue_IncrementAndDecrement() { userID := int64(20) waitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) @@ -312,9 +353,11 @@ func (s *ConcurrencyCacheSuite) TestAccountWaitQueue_IncrementAndDecrement() { func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { accountID := int64(901) userID := int64(902) + apiKeyID := int64(903) unindexedAccountID := int64(1901) accountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID) userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, userID) + apiKeyKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) unindexedAccountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, unindexedAccountID) userWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID) @@ -333,6 +376,10 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { require.NoError(s.T(), s.rdb.ZAdd(s.ctx, unindexedAccountKey, redis.Z{Score: float64(now), Member: "oldproc-unindexed"}, ).Err()) + require.NoError(s.T(), s.rdb.ZAdd(s.ctx, apiKeyKey, + redis.Z{Score: float64(now), Member: "oldproc-3"}, + redis.Z{Score: float64(now), Member: "keep-3"}, + ).Err()) require.NoError(s.T(), s.rdb.Set(s.ctx, userWaitKey, 3, time.Minute).Err()) require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 2, time.Minute).Err()) require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedAccountWaitKey, 2, time.Minute).Err()) @@ -355,6 +402,11 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { require.NoError(s.T(), err) require.Equal(s.T(), []string{"keep-2"}, userMembers) + // API Key 槽位(stats-only)不在启动清理范围内,靠分数裁剪与 key TTL 自愈。 + apiKeyMembers, err := s.rdb.ZRange(s.ctx, apiKeyKey, 0, -1).Result() + require.NoError(s.T(), err) + require.ElementsMatch(s.T(), []string{"keep-3", "oldproc-3"}, apiKeyMembers) + _, err = s.rdb.Get(s.ctx, userWaitKey).Result() require.True(s.T(), errors.Is(err, redis.Nil)) diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 4e839b6a12..a4e173006e 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -466,6 +466,49 @@ func (r *groupRepository) ListActive(ctx context.Context) ([]service.Group, erro return outGroups, nil } +func (r *groupRepository) ListActiveIDs(ctx context.Context) ([]int64, error) { + if r.sql != nil { + rows, err := r.sql.QueryContext(ctx, ` + SELECT id + FROM groups + WHERE status = $1 + AND deleted_at IS NULL + ORDER BY sort_order ASC, id ASC + `, service.StatusActive) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + ids := make([]int64, 0) + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + ids = append(ids, id) + } + if err := rows.Err(); err != nil { + return nil, err + } + return ids, nil + } + + groups, err := r.client.Group.Query(). + Where(group.StatusEQ(service.StatusActive)). + Select(group.FieldID). + Order(dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)). + All(ctx) + if err != nil { + return nil, err + } + ids := make([]int64, 0, len(groups)) + for i := range groups { + ids = append(ids, groups[i].ID) + } + return ids, nil +} + func (r *groupRepository) ListActiveByPlatform(ctx context.Context, platform string) ([]service.Group, error) { groups, err := r.client.Group.Query(). Where(group.StatusEQ(service.StatusActive), group.PlatformEQ(platform)). diff --git a/backend/internal/repository/ops_error_where_test.go b/backend/internal/repository/ops_error_where_test.go index 5b9d7ab1c3..c997865ba4 100644 --- a/backend/internal/repository/ops_error_where_test.go +++ b/backend/internal/repository/ops_error_where_test.go @@ -85,10 +85,21 @@ func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) { t.Fatalf("default filter must still include the status >= 400 guard for non-cyber rows\nfull: %s", where) } - // phase=upstream skips the status guard entirely — exemption is irrelevant there. + // phase=upstream WITHOUT the recovered-upstream opt-in keeps the status guard: + // request-error list endpoints filter by phase=upstream as a plain condition. whereUpstream, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "upstream"}) - if strings.Contains(whereUpstream, "status_code") { - t.Fatalf("upstream phase filter must not add any status_code clause\nfull: %s", whereUpstream) + if !strings.Contains(whereUpstream, "COALESCE(e.status_code, 0) >= 400") { + t.Fatalf("upstream phase without IncludeRecoveredUpstream must keep the status guard\nfull: %s", whereUpstream) + } + if !strings.Contains(whereUpstream, "e.error_phase = $") { + t.Fatalf("upstream phase filter must emit the error_phase condition\nfull: %s", whereUpstream) + } + + // phase=upstream WITH IncludeRecoveredUpstream (ops 上游列表) skips the guard, + // exposing recovered (<400) upstream rows. + whereRecovered, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "upstream", IncludeRecoveredUpstream: true}) + if strings.Contains(whereRecovered, "status_code") { + t.Fatalf("upstream phase with IncludeRecoveredUpstream must not add any status_code clause\nfull: %s", whereRecovered) } } diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go index 9923c08d99..2129a451c4 100644 --- a/backend/internal/repository/ops_repo.go +++ b/backend/internal/repository/ops_repo.go @@ -177,6 +177,37 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any { } } +// opsErrorLogsOrderBy builds the ORDER BY clause from a whitelist, mirroring +// usageLogOrderBy semantics. Unknown SortBy falls back to created_at; e.id is +// always appended as tiebreaker for stable pagination. +func opsErrorLogsOrderBy(filter *service.OpsErrorLogFilter) string { + sortBy := "" + sortOrder := "" + if filter != nil { + sortBy = strings.ToLower(strings.TrimSpace(filter.SortBy)) + sortOrder = strings.ToLower(strings.TrimSpace(filter.SortOrder)) + } + + var column string + switch sortBy { + case "model": + column = "COALESCE(NULLIF(TRIM(e.requested_model), ''), e.model)" + case "status_code": + // 与展示列/过滤保持同义:列表展示 COALESCE(upstream_status_code, status_code, 0), + // status_code 过滤也用同一表达式,故排序必须一致——否则 recovered upstream 行 + //(status_code<400 但展示上游 5xx)排序键与显示值/分页切分不符。 + column = "COALESCE(e.upstream_status_code, e.status_code, 0)" + default: + column = "e.created_at" + } + + dir := "DESC" + if sortOrder == "asc" { + dir = "ASC" + } + return fmt.Sprintf("%s %s, e.id %s", column, dir, dir) +} + func (r *opsRepository) ListErrorLogs(ctx context.Context, filter *service.OpsErrorLogFilter) (*service.OpsErrorLogList, error) { if r == nil || r.db == nil { return nil, fmt.Errorf("nil ops repository") @@ -233,25 +264,29 @@ SELECT COALESCE(a.name, ''), e.group_id, COALESCE(g.name, ''), - CASE WHEN e.client_ip IS NULL THEN NULL ELSE e.client_ip::text END, + CASE WHEN e.client_ip IS NULL THEN NULL ELSE host(e.client_ip) END, COALESCE(e.request_path, ''), e.stream, COALESCE(e.inbound_endpoint, ''), COALESCE(e.upstream_endpoint, ''), COALESCE(e.requested_model, ''), COALESCE(e.upstream_model, ''), + COALESCE(e.user_agent, ''), e.request_type, COALESCE(ak.name, ''), ak.deleted_at, - COALESCE(e.deleted_key_name, '') + COALESCE(e.deleted_key_name, ''), + e.deleted_key_owner_user_id, + COALESCE(du.email, '') FROM ops_error_logs e LEFT JOIN accounts a ON e.account_id = a.id LEFT JOIN groups g ON e.group_id = g.id LEFT JOIN users u ON e.user_id = u.id LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id +LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id LEFT JOIN api_keys ak ON ak.id = e.api_key_id ` + where + ` -ORDER BY e.created_at DESC +ORDER BY ` + opsErrorLogsOrderBy(filter) + ` LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) rows, err := r.db.QueryContext(ctx, selectSQL, argsWithLimit...) @@ -279,6 +314,8 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) var apiKeyName string var apiKeyDeletedAt sql.NullTime var deletedKeyName string + var deletedKeyOwnerID sql.NullInt64 + var deletedKeyOwnerEmail string if err := rows.Scan( &item.ID, &item.CreatedAt, @@ -311,10 +348,13 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) &item.UpstreamEndpoint, &item.RequestedModel, &item.UpstreamModel, + &item.UserAgent, &requestType, &apiKeyName, &apiKeyDeletedAt, &deletedKeyName, + &deletedKeyOwnerID, + &deletedKeyOwnerEmail, ); err != nil { return nil, err } @@ -364,6 +404,12 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) } // 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。 item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "") + // 已删除 KEY 所有者快照:认证失败行 user_id 为空,列表用户列以此回退。 + if deletedKeyOwnerID.Valid { + v := deletedKeyOwnerID.Int64 + item.DeletedKeyOwnerUserID = &v + item.DeletedKeyOwnerEmail = deletedKeyOwnerEmail + } out = append(out, &item) } if err := rows.Err(); err != nil { @@ -417,7 +463,7 @@ SELECT COALESCE(a.name, ''), e.group_id, COALESCE(g.name, ''), - CASE WHEN e.client_ip IS NULL THEN NULL ELSE e.client_ip::text END, + CASE WHEN e.client_ip IS NULL THEN NULL ELSE host(e.client_ip) END, COALESCE(e.request_path, ''), e.stream, COALESCE(e.inbound_endpoint, ''), @@ -927,12 +973,14 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) { if filter != nil { resolvedFilter = filter.Resolved } - // Keep list endpoints scoped to client errors unless explicitly filtering upstream phase. + // Keep list endpoints scoped to client errors unless the caller explicitly opts + // into recovered upstream rows (Phase=="upstream" + IncludeRecoveredUpstream, + // ops 专用上游列表)。请求错误语义的端点即便过滤 phase=upstream 也保留该守卫。 // cyber_policy is exempt from the status >= 400 guard: streaming cyber hits arrive with // status 200 (the SSE stream opened successfully before upstream returned response.failed), // but they are always client-visible blocked requests that belong in admin + user error // lists. Without the exemption the entire streaming-path cyber sink would be invisible. - if phaseFilter != "upstream" { + if phaseFilter != "upstream" || filter == nil || !filter.IncludeRecoveredUpstream { clauses = append(clauses, "(COALESCE(e.status_code, 0) >= 400 OR e.error_type = 'cyber_policy')") } diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index c508f09c71..c8e1fe14e0 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -518,7 +518,7 @@ func filterSchedulerCredentials(credentials map[string]any) map[string]any { if len(credentials) == 0 { return nil } - keys := []string{"model_mapping", "compact_model_mapping", "api_key", "project_id", "oauth_type"} + keys := []string{"model_mapping", "compact_model_mapping", "api_key", "project_id", "oauth_type", "plan_type"} filtered := make(map[string]any) for _, key := range keys { if value, ok := credentials[key]; ok && value != nil { diff --git a/backend/internal/repository/scheduler_cache_test.go b/backend/internal/repository/scheduler_cache_test.go new file mode 100644 index 0000000000..f438b2f456 --- /dev/null +++ b/backend/internal/repository/scheduler_cache_test.go @@ -0,0 +1,37 @@ +package repository + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestFilterSchedulerCredentialsKeepsSubscriptionPlanType(t *testing.T) { + filtered := filterSchedulerCredentials(map[string]any{ + "plan_type": "plus", + "access_token": "secret-access-token", + "refresh_token": "secret-refresh-token", + }) + + require.Equal(t, "plus", filtered["plan_type"]) + require.NotContains(t, filtered, "access_token") + require.NotContains(t, filtered, "refresh_token") +} + +func TestSchedulerMetadataAccountKeepsOpenAISubscriptionIdentity(t *testing.T) { + account := service.Account{ + ID: 24, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "plan_type": "plus", + "access_token": "secret-access-token", + }, + } + + metadata := buildSchedulerMetadataAccount(account) + + require.True(t, metadata.IsOpenAIChatGPTSubscription()) + require.Empty(t, metadata.GetCredential("access_token")) +} diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index 885e63f9fd..24c648b0a5 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -372,12 +372,13 @@ func (r *usageLogRepository) CreateBestEffort(ctx context.Context, log *service. } } + // 队列满时阻塞等待而非立即丢弃:批处理器持续排空队列,短暂等待即可入队。 + // 立即丢弃会造成“已扣费但无 usage_log”的永久数据缺口(issue #3656); + // 阻塞上限由调用方 ctx 期限约束,超时后由上层同步兜底。 select { case r.bestEffortBatchCh <- req: case <-ctx.Done(): return service.MarkUsageLogCreateDropped(ctx.Err()) - default: - return service.MarkUsageLogCreateDropped(errors.New("usage log best-effort queue full")) } select { @@ -493,12 +494,12 @@ func (r *usageLogRepository) createBatched(ctx context.Context, log *service.Usa resultCh: make(chan usageLogCreateResult, 1), } + // 队列满时阻塞等待而非立即报错:本路径是 best-effort 丢弃后的最后兜底, + // 立即失败会让日志永久丢失;阻塞上限由调用方 ctx 期限约束。 select { case r.createBatchCh <- req: case <-ctx.Done(): return false, service.MarkUsageLogCreateNotPersisted(ctx.Err()) - default: - return false, service.MarkUsageLogCreateNotPersisted(errors.New("usage log create batch queue full")) } select { @@ -520,22 +521,28 @@ func (r *usageLogRepository) createBatched(ctx context.Context, log *service.Usa } func (r *usageLogRepository) ensureCreateBatcher() { - if r == nil || r.db == nil || r.createBatchCh != nil { + if r == nil || r.db == nil { return } + // nil 检查必须在 Once 内部:在外层做无同步快路径读会与 Once 内的写构成数据竞争。 r.createBatchOnce.Do(func() { - r.createBatchCh = make(chan usageLogCreateRequest, usageLogCreateBatchQueueCap) - go r.runCreateBatcher(r.db) + if r.createBatchCh == nil { + r.createBatchCh = make(chan usageLogCreateRequest, usageLogCreateBatchQueueCap) + go r.runCreateBatcher(r.db) + } }) } func (r *usageLogRepository) ensureBestEffortBatcher() { - if r == nil || r.db == nil || r.bestEffortBatchCh != nil { + if r == nil || r.db == nil { return } + // 同 ensureCreateBatcher:nil 检查放在 Once 内部以避免数据竞争。 r.bestEffortBatchOnce.Do(func() { - r.bestEffortBatchCh = make(chan usageLogBestEffortRequest, usageLogBestEffortBatchQueueCap) - go r.runBestEffortBatcher(r.db) + if r.bestEffortBatchCh == nil { + r.bestEffortBatchCh = make(chan usageLogBestEffortRequest, usageLogBestEffortBatchQueueCap) + go r.runBestEffortBatcher(r.db) + } }) } diff --git a/backend/internal/repository/usage_log_repo_integration_test.go b/backend/internal/repository/usage_log_repo_integration_test.go index ed3050d89c..b43c6d56ef 100644 --- a/backend/internal/repository/usage_log_repo_integration_test.go +++ b/backend/internal/repository/usage_log_repo_integration_test.go @@ -288,21 +288,21 @@ func TestUsageLogRepositoryCreateBestEffort_BatchPathDuplicateRequestID(t *testi }, 3*time.Second, 20*time.Millisecond) } -func TestUsageLogRepositoryCreateBestEffort_QueueFullReturnsDropped(t *testing.T) { - ctx := context.Background() +func TestUsageLogRepositoryCreateBestEffort_QueueFullBlocksUntilCtxDeadline(t *testing.T) { + // 队列满时不再立即丢弃:阻塞等待入队,直到调用方 ctx 到期才标记 dropped(issue #3656)。 client := testEntClient(t) repo := newUsageLogRepositoryWithSQL(client, integrationDB) repo.bestEffortBatchCh = make(chan usageLogBestEffortRequest, 1) repo.bestEffortBatchCh <- usageLogBestEffortRequest{} - user := mustCreateUser(t, client, &service.User{Email: fmt.Sprintf("usage-best-effort-full-%d@example.com", time.Now().UnixNano())}) - apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-usage-best-effort-full-" + uuid.NewString(), Name: "k"}) - account := mustCreateAccount(t, client, &service.Account{Name: "acc-usage-best-effort-full-" + uuid.NewString()}) + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + start := time.Now() err := repo.CreateBestEffort(ctx, &service.UsageLog{ - UserID: user.ID, - APIKeyID: apiKey.ID, - AccountID: account.ID, + UserID: 1, + APIKeyID: 2, + AccountID: 3, RequestID: uuid.NewString(), Model: "claude-3", InputTokens: 10, @@ -314,6 +314,40 @@ func TestUsageLogRepositoryCreateBestEffort_QueueFullReturnsDropped(t *testing.T require.Error(t, err) require.True(t, service.IsUsageLogCreateDropped(err)) + require.GreaterOrEqual(t, time.Since(start), 150*time.Millisecond) +} + +func TestUsageLogRepositoryCreateBestEffort_QueueFullWaitsForDrain(t *testing.T) { + // 队列满但批处理器随后排空时,阻塞的入队应成功完成而非丢弃。 + client := testEntClient(t) + repo := newUsageLogRepositoryWithSQL(client, integrationDB) + repo.bestEffortBatchCh = make(chan usageLogBestEffortRequest, 1) + repo.bestEffortBatchCh <- usageLogBestEffortRequest{} + + go func() { + time.Sleep(100 * time.Millisecond) + <-repo.bestEffortBatchCh // 排空占位请求,为阻塞中的入队腾出空间 + req := <-repo.bestEffortBatchCh + sendUsageLogBestEffortResult(req.resultCh, nil) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + err := repo.CreateBestEffort(ctx, &service.UsageLog{ + UserID: 1, + APIKeyID: 2, + AccountID: 3, + RequestID: uuid.NewString(), + Model: "claude-3", + InputTokens: 10, + OutputTokens: 20, + TotalCost: 0.5, + ActualCost: 0.5, + CreatedAt: time.Now().UTC(), + }) + + require.NoError(t, err) } func TestUsageLogRepositoryCreate_BatchPathCanceledContextMarksNotPersisted(t *testing.T) { @@ -346,7 +380,7 @@ func TestUsageLogRepositoryCreate_BatchPathCanceledContextMarksNotPersisted(t *t } func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing.T) { - ctx := context.Background() + // 队列满时阻塞等待入队,直到调用方 ctx 到期才标记 not persisted(issue #3656)。 client := testEntClient(t) repo := newUsageLogRepositoryWithSQL(client, integrationDB) repo.createBatchCh = make(chan usageLogCreateRequest, 1) @@ -356,6 +390,10 @@ func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-usage-create-full-" + uuid.NewString(), Name: "k"}) account := mustCreateAccount(t, client, &service.Account{Name: "acc-usage-create-full-" + uuid.NewString()}) + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + start := time.Now() inserted, err := repo.Create(ctx, &service.UsageLog{ UserID: user.ID, APIKeyID: apiKey.ID, @@ -372,6 +410,7 @@ func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing require.False(t, inserted) require.Error(t, err) require.True(t, service.IsUsageLogCreateNotPersisted(err)) + require.GreaterOrEqual(t, time.Since(start), 150*time.Millisecond) } func TestUsageLogRepositoryCreate_BatchPathCanceledAfterQueueMarksNotPersisted(t *testing.T) { diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 26f976c87b..e3e68bb6e7 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -233,6 +233,7 @@ func TestAPIContracts(t *testing.T) { "ip_whitelist": null, "ip_blacklist": null, "last_used_at": null, + "current_concurrency": 0, "quota": 0, "quota_used": 0, "rate_limit_5h": 0, @@ -282,6 +283,7 @@ func TestAPIContracts(t *testing.T) { "ip_whitelist": null, "ip_blacklist": null, "last_used_at": null, + "current_concurrency": 0, "quota": 0, "quota_used": 0, "rate_limit_5h": 0, @@ -661,15 +663,17 @@ func TestAPIContracts(t *testing.T) { service.SettingKeyTableDefaultPageSize: "20", service.SettingKeyTablePageSizeOptions: "[10,20,50,100]", - service.SettingKeyOpsMonitoringEnabled: "false", - service.SettingKeyOpsRealtimeMonitoringEnabled: "true", - service.SettingKeyOpsQueryModeDefault: "auto", - service.SettingKeyOpsMetricsIntervalSeconds: "60", - service.SettingPaymentVisibleMethodAlipaySource: service.VisibleMethodSourceEasyPayAlipay, - service.SettingPaymentVisibleMethodWxpaySource: service.VisibleMethodSourceOfficialWechat, - service.SettingPaymentVisibleMethodAlipayEnabled: "true", - service.SettingPaymentVisibleMethodWxpayEnabled: "false", - "openai_advanced_scheduler_enabled": "true", + service.SettingKeyOpsMonitoringEnabled: "false", + service.SettingKeyOpsRealtimeMonitoringEnabled: "true", + service.SettingKeyOpsQueryModeDefault: "auto", + service.SettingKeyOpsMetricsIntervalSeconds: "60", + service.SettingPaymentVisibleMethodAlipaySource: service.VisibleMethodSourceEasyPayAlipay, + service.SettingPaymentVisibleMethodWxpaySource: service.VisibleMethodSourceOfficialWechat, + service.SettingPaymentVisibleMethodAlipayEnabled: "true", + service.SettingPaymentVisibleMethodWxpayEnabled: "false", + "openai_advanced_scheduler_enabled": "true", + service.SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled: "false", + service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled: "false", }) }, method: http.MethodGet, @@ -861,6 +865,28 @@ func TestAPIContracts(t *testing.T) { "payment_visible_method_alipay_enabled": true, "payment_visible_method_wxpay_enabled": false, "openai_advanced_scheduler_enabled": true, + "openai_advanced_scheduler_sticky_weighted_enabled": false, + "openai_advanced_scheduler_subscription_priority_enabled": false, + "openai_advanced_scheduler_lb_top_k": "", + "openai_advanced_scheduler_weight_priority": "", + "openai_advanced_scheduler_weight_load": "", + "openai_advanced_scheduler_weight_queue": "", + "openai_advanced_scheduler_weight_error_rate": "", + "openai_advanced_scheduler_weight_ttft": "", + "openai_advanced_scheduler_weight_reset": "", + "openai_advanced_scheduler_weight_quota_headroom": "", + "openai_advanced_scheduler_weight_previous_response": "", + "openai_advanced_scheduler_weight_session_sticky": "", + "openai_advanced_scheduler_effective_lb_top_k": "7", + "openai_advanced_scheduler_effective_weight_priority": "1", + "openai_advanced_scheduler_effective_weight_load": "1", + "openai_advanced_scheduler_effective_weight_queue": "0.7", + "openai_advanced_scheduler_effective_weight_error_rate": "0.8", + "openai_advanced_scheduler_effective_weight_ttft": "0.5", + "openai_advanced_scheduler_effective_weight_reset": "0", + "openai_advanced_scheduler_effective_weight_quota_headroom": "0", + "openai_advanced_scheduler_effective_weight_previous_response": "5", + "openai_advanced_scheduler_effective_weight_session_sticky": "3", "openai_codex_user_agent": "", "openai_fast_policy_settings": { "rules": [] @@ -875,6 +901,7 @@ func TestAPIContracts(t *testing.T) { "payment_max_pending_orders": 0, "payment_balance_disabled": false, "payment_balance_recharge_multiplier": 0, + "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, "payment_load_balance_strategy": "", "payment_product_name_prefix": "", @@ -1110,6 +1137,28 @@ func TestAPIContracts(t *testing.T) { "payment_visible_method_alipay_enabled": false, "payment_visible_method_wxpay_enabled": false, "openai_advanced_scheduler_enabled": false, + "openai_advanced_scheduler_sticky_weighted_enabled": false, + "openai_advanced_scheduler_subscription_priority_enabled": false, + "openai_advanced_scheduler_lb_top_k": "", + "openai_advanced_scheduler_weight_priority": "", + "openai_advanced_scheduler_weight_load": "", + "openai_advanced_scheduler_weight_queue": "", + "openai_advanced_scheduler_weight_error_rate": "", + "openai_advanced_scheduler_weight_ttft": "", + "openai_advanced_scheduler_weight_reset": "", + "openai_advanced_scheduler_weight_quota_headroom": "", + "openai_advanced_scheduler_weight_previous_response": "", + "openai_advanced_scheduler_weight_session_sticky": "", + "openai_advanced_scheduler_effective_lb_top_k": "7", + "openai_advanced_scheduler_effective_weight_priority": "1", + "openai_advanced_scheduler_effective_weight_load": "1", + "openai_advanced_scheduler_effective_weight_queue": "0.7", + "openai_advanced_scheduler_effective_weight_error_rate": "0.8", + "openai_advanced_scheduler_effective_weight_ttft": "0.5", + "openai_advanced_scheduler_effective_weight_reset": "0", + "openai_advanced_scheduler_effective_weight_quota_headroom": "0", + "openai_advanced_scheduler_effective_weight_previous_response": "5", + "openai_advanced_scheduler_effective_weight_session_sticky": "3", "openai_codex_user_agent": "", "openai_fast_policy_settings": { "rules": [] @@ -1123,6 +1172,7 @@ func TestAPIContracts(t *testing.T) { "payment_enabled_types": null, "payment_balance_disabled": false, "payment_balance_recharge_multiplier": 0, + "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, "payment_load_balance_strategy": "", "payment_product_name_prefix": "", @@ -1690,6 +1740,10 @@ func (s *stubAccountRepo) List(ctx context.Context, params pagination.Pagination return nil, nil, errors.New("not implemented") } +func (s *stubAccountRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]service.Account, error) { + return nil, nil +} + func (s *stubAccountRepo) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) { return nil, nil, errors.New("not implemented") } diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 7d5e70a15b..ae25bf387d 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -70,6 +70,14 @@ type Account struct { modelMappingCacheRawPtr uintptr modelMappingCacheRawLen int modelMappingCacheRawSig uint64 + + // header_overrides 热路径缓存(非持久化字段,同 model_mapping 缓存先例) + headerOverrideCache map[string]string + headerOverrideCacheReady bool + headerOverrideCacheCredentialsPtr uintptr + headerOverrideCacheRawPtr uintptr + headerOverrideCacheRawLen int + headerOverrideCacheRawSig uint64 } type OpenAIEndpointCapability string @@ -580,6 +588,7 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri "gemini-3.1-pro-high", "gemini-3.1-pro-low", }) + applyAntigravityGemini31ProAliases(result) } return result } @@ -646,6 +655,61 @@ func ensureAntigravityDefaultPassthroughs(mapping map[string]string, models []st } } +func applyAntigravityGemini31ProAliases(mapping map[string]string) { + target := strings.TrimSpace(mapping[domain.AntigravityGemini31ProAgentModel]) + if target == "" { + return + } + + aliases := []struct { + model string + legacyTargets map[string]struct{} + }{ + { + model: "gemini-3.1-pro", + legacyTargets: map[string]struct{}{ + "gemini-3.1-pro": {}, + }, + }, + { + model: "gemini-3.1-pro-high", + legacyTargets: map[string]struct{}{ + "gemini-3.1-pro-high": {}, + }, + }, + { + model: "gemini-3.1-pro-preview", + legacyTargets: map[string]struct{}{ + "gemini-3.1-pro-preview": {}, + "gemini-3.1-pro-high": {}, + }, + }, + } + + for _, alias := range aliases { + current, exists := mapping[alias.model] + if exists { + if _, legacy := alias.legacyTargets[current]; legacy { + mapping[alias.model] = target + } + continue + } + if mappingHasWildcardForModel(mapping, alias.model) { + continue + } + mapping[alias.model] = target + } +} + +func mappingHasWildcardForModel(mapping map[string]string, model string) bool { + for pattern := range mapping { + if matchWildcard(pattern, model) { + return true + } + } + return false +} + func normalizeRequestedModelForLookup(platform, requestedModel string) string { trimmed := strings.TrimSpace(requestedModel) if trimmed == "" { @@ -1126,6 +1190,18 @@ func (a *Account) IsOpenAIOAuth() bool { return a.IsOpenAI() && a.Type == AccountTypeOAuth } +func (a *Account) IsOpenAIChatGPTSubscription() bool { + if !a.IsOpenAIOAuth() { + return false + } + switch strings.ToLower(strings.TrimSpace(a.GetCredential("plan_type"))) { + case "", "free", "abnormal": + return false + default: + return true + } +} + func (a *Account) IsOpenAIPersonalAccessToken() bool { if !a.IsOpenAIOAuth() { return false diff --git a/backend/internal/service/account_header_override.go b/backend/internal/service/account_header_override.go new file mode 100644 index 0000000000..8882bbef91 --- /dev/null +++ b/backend/internal/service/account_header_override.go @@ -0,0 +1,280 @@ +package service + +import ( + "net/http" + "strings" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + + "golang.org/x/net/http/httpguts" +) + +// 请求头覆写(header override):仅对 Anthropic / OpenAI 平台的 api_key 账号生效。 +// 管理员在账号上配置一组 header name -> value,转发到上游前用配置值覆盖同名请求头 +// (匹配不区分大小写);value 为空的条目视为"未填写",不参与覆盖。 +const ( + credKeyHeaderOverrideEnabled = "header_override_enabled" + credKeyHeaderOverrides = "header_overrides" + + maxHeaderOverrideEntries = 64 + maxHeaderOverrideNameLength = 200 + maxHeaderOverrideValueLength = 8192 +) + +// headerOverrideBlockedNames 禁止覆写的请求头(小写)。 +// - 连接控制/逐跳头:由 HTTP 栈管理,覆写会破坏请求传输; +// - host/content-length:由 Go 的 Request.Host / ContentLength 字段管理,header 覆写不生效或产生冲突; +// - content-type:承载报文框架信息(multipart boundary 为每请求随机值),静态覆写必然与 body 不匹配; +// - authorization/x-api-key/cookie 等:上游认证头由账号凭据统一注入,禁止通过覆写篡改或重新引入; +// - accept-encoding:强制压缩会破坏网关对上游流式响应(SSE/usage)的解析; +// - sec-websocket-*:WebSocket 握手头由拨号器管理(OpenAI WS 模式); +// - session_id/x-claude-code-session-id 等:逐请求会话隔离头,固定值会造成会话串扰。 +var headerOverrideBlockedNames = map[string]struct{}{ + "host": {}, + "content-length": {}, + "content-type": {}, + "transfer-encoding": {}, + "connection": {}, + "keep-alive": {}, + "proxy-authenticate": {}, + "proxy-authorization": {}, + "proxy-connection": {}, + "te": {}, + "trailer": {}, + "upgrade": {}, + "authorization": {}, + "x-api-key": {}, + "x-goog-api-key": {}, + "cookie": {}, + "accept-encoding": {}, + "sec-websocket-key": {}, + "sec-websocket-version": {}, + "sec-websocket-extensions": {}, + "sec-websocket-protocol": {}, + "sec-websocket-accept": {}, + "session_id": {}, + "conversation_id": {}, + "x-codex-turn-state": {}, + "x-codex-turn-metadata": {}, + "chatgpt-account-id": {}, + "x-claude-code-session-id": {}, + "x-client-request-id": {}, +} + +func isHeaderOverrideBlockedName(lowerName string) bool { + _, blocked := headerOverrideBlockedNames[lowerName] + return blocked +} + +// IsHeaderOverrideEligible 报告账号类型是否支持请求头覆写。 +// 目前仅开放 Anthropic / OpenAI 两个平台的 api_key 账号。 +func (a *Account) IsHeaderOverrideEligible() bool { + if a == nil || a.Type != AccountTypeAPIKey { + return false + } + return a.Platform == PlatformAnthropic || a.Platform == PlatformOpenAI +} + +// IsHeaderOverrideEnabled 报告账号是否启用了请求头覆写。 +func (a *Account) IsHeaderOverrideEnabled() bool { + if !a.IsHeaderOverrideEligible() || a.Credentials == nil { + return false + } + enabled, ok := a.Credentials[credKeyHeaderOverrideEnabled].(bool) + return ok && enabled +} + +// GetHeaderOverrides 返回生效的请求头覆写表(key 统一小写)。 +// 未启用、不符合平台/类型条件或配置为空时返回 nil。 +// 空 value 的条目(模板占位)与非法/禁止的 header 名会被跳过。 +// 结果带热路径缓存(同 GetModelMapping 先例):同一 credentials 映射在 +// 一次请求 / 一条 WS 会话内的多次调用只做一次解析与校验。 +func (a *Account) GetHeaderOverrides() map[string]string { + if !a.IsHeaderOverrideEnabled() { + return nil + } + rawMapping, rawIsAnyMap := a.Credentials[credKeyHeaderOverrides].(map[string]any) + if !rawIsAnyMap { + // 非 JSON 反序列化产物(如直接注入的 map[string]string):直接解析,不缓存 + return resolveHeaderOverrides(stringMappingFromRaw(a.Credentials[credKeyHeaderOverrides])) + } + + credentialsPtr := mapPtr(a.Credentials) + rawPtr := mapPtr(rawMapping) + rawLen := len(rawMapping) + rawSig := uint64(0) + rawSigReady := false + + if a.headerOverrideCacheReady && + a.headerOverrideCacheCredentialsPtr == credentialsPtr && + a.headerOverrideCacheRawPtr == rawPtr && + a.headerOverrideCacheRawLen == rawLen { + rawSig = modelMappingSignature(rawMapping) + rawSigReady = true + if a.headerOverrideCacheRawSig == rawSig { + return a.headerOverrideCache + } + } + + overrides := resolveHeaderOverrides(stringMappingFromRaw(rawMapping)) + if !rawSigReady { + rawSig = modelMappingSignature(rawMapping) + } + + a.headerOverrideCache = overrides + a.headerOverrideCacheReady = true + a.headerOverrideCacheCredentialsPtr = credentialsPtr + a.headerOverrideCacheRawPtr = rawPtr + a.headerOverrideCacheRawLen = rawLen + a.headerOverrideCacheRawSig = rawSig + return overrides +} + +// resolveHeaderOverrides 解析并防御性过滤原始覆写表:保存路径已做校验, +// 这里兜底未经 Normalize 落库的数据(含名单扩充前保存的旧配置),非法条目直接跳过。 +func resolveHeaderOverrides(raw map[string]string) map[string]string { + if len(raw) == 0 { + return nil + } + result := make(map[string]string, len(raw)) + for name, value := range raw { + lowerName, value, err := normalizeHeaderOverrideEntry(name, value) + if err != nil || lowerName == "" || value == "" { + continue + } + result[lowerName] = value + } + if len(result) == 0 { + return nil + } + return result +} + +// HeaderOverrideValue 返回指定 header(小写名)的生效覆写值。 +// 供转发链路在 header 写入前感知覆写结果(如 anthropic-beta 需要参与 body 净化)。 +func (a *Account) HeaderOverrideValue(lowerName string) (string, bool) { + value, ok := a.GetHeaderOverrides()[lowerName] + return value, ok +} + +// ApplyHeaderOverrides 将账号配置的请求头覆写应用到出站请求头。 +// 对每个覆写条目:先删除所有大小写变体(转发链路会以 wire casing 直接写入 map, +// 可能存在非 canonical key),再按已知 wire casing 写入,避免产生重复头。 +// 账号未启用或不符合条件时为 no-op,可安全地在 OAuth/api_key 共用的构建器中调用。 +func (a *Account) ApplyHeaderOverrides(h http.Header) { + if h == nil { + return + } + overrides := a.GetHeaderOverrides() + if len(overrides) == 0 { + return + } + // 覆写名两两不同(大小写不敏感)且各自只操作同名键,应用顺序不影响结果。 + // 全量 EqualFold 扫描兜底删除任意 casing 的既有键:透传链路可能保留客户端 + // 原始 casing,非 canonical/wire casing 的键 deleteHeaderAllForms 覆盖不到。 + for name, value := range overrides { + for existing := range h { + if strings.EqualFold(existing, name) { + delete(h, existing) + } + } + h[resolveWireCasing(name)] = []string{value} + } +} + +// NormalizeHeaderOverrideCredentials 校验并原地规范化 credentials 中的请求头覆写字段。 +// 供账号创建/更新/批量更新的保存路径调用;credentials 未携带相关字段时为 no-op。 +// 规范化内容:header 名转小写并去除首尾空白,value 去除首尾空白,丢弃名和值均为空的条目。 +func NormalizeHeaderOverrideCredentials(credentials map[string]any) error { + if credentials == nil { + return nil + } + if raw, ok := credentials[credKeyHeaderOverrideEnabled]; ok && raw != nil { + if _, isBool := raw.(bool); !isBool { + return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_override_enabled must be a boolean") + } + } + raw, ok := credentials[credKeyHeaderOverrides] + if !ok || raw == nil { + return nil + } + + var entries map[string]any + switch m := raw.(type) { + case map[string]any: + entries = m + case map[string]string: + entries = make(map[string]any, len(m)) + for k, v := range m { + entries[k] = v + } + default: + return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_overrides must be an object of header name to string value") + } + + if len(entries) > maxHeaderOverrideEntries { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_overrides supports at most %d entries", maxHeaderOverrideEntries) + } + + normalized := make(map[string]any, len(entries)) + for name, rawValue := range entries { + value, isString := rawValue.(string) + if !isString { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q value must be a string", name) + } + lowerName, value, err := normalizeHeaderOverrideEntry(name, value) + if err != nil { + return err + } + if lowerName == "" { + continue // 丢弃完全为空的占位行 + } + if _, dup := normalized[lowerName]; dup { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "duplicate header name %q (matching is case-insensitive)", lowerName) + } + normalized[lowerName] = value + } + credentials[credKeyHeaderOverrides] = normalized + return nil +} + +// normalizeHeaderOverrideEntry 校验并规范化单个覆写条目,保存路径(Normalize,err → 400) +// 与应用路径(resolveHeaderOverrides,err → 跳过)共用同一套规则,避免两处校验漂移。 +// 名和值均为空表示空占位行,返回 ("", "", nil);空 value 的具名条目合法(模板占位)。 +func normalizeHeaderOverrideEntry(name, value string) (string, string, error) { + lowerName := strings.ToLower(strings.TrimSpace(name)) + value = strings.TrimSpace(value) + if lowerName == "" { + if value == "" { + return "", "", nil + } + return "", "", infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header name must not be empty") + } + if len(lowerName) > maxHeaderOverrideNameLength { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header name %q exceeds %d characters", lowerName, maxHeaderOverrideNameLength) + } + if !httpguts.ValidHeaderFieldName(lowerName) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "invalid header name %q", lowerName) + } + if isHeaderOverrideBlockedName(lowerName) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q is not allowed to be overridden", lowerName) + } + if len(value) > maxHeaderOverrideValueLength { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q value exceeds %d characters", lowerName, maxHeaderOverrideValueLength) + } + if !httpguts.ValidHeaderFieldValue(value) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q has an invalid value", lowerName) + } + return lowerName, value, nil +} diff --git a/backend/internal/service/account_header_override_test.go b/backend/internal/service/account_header_override_test.go new file mode 100644 index 0000000000..c89b5e0587 --- /dev/null +++ b/backend/internal/service/account_header_override_test.go @@ -0,0 +1,339 @@ +//go:build unit + +package service + +import ( + "net/http" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func headerOverrideTestAccount(platform, accountType string, credentials map[string]any) *Account { + return &Account{ + Platform: platform, + Type: accountType, + Credentials: credentials, + } +} + +func TestIsHeaderOverrideEligible(t *testing.T) { + tests := []struct { + name string + platform string + accType string + want bool + }{ + {"anthropic apikey", PlatformAnthropic, AccountTypeAPIKey, true}, + {"openai apikey", PlatformOpenAI, AccountTypeAPIKey, true}, + {"anthropic oauth", PlatformAnthropic, AccountTypeOAuth, false}, + {"openai oauth", PlatformOpenAI, AccountTypeOAuth, false}, + {"gemini apikey", PlatformGemini, AccountTypeAPIKey, false}, + {"grok apikey", PlatformGrok, AccountTypeAPIKey, false}, + {"antigravity apikey", PlatformAntigravity, AccountTypeAPIKey, false}, + {"anthropic bedrock", PlatformAnthropic, AccountTypeBedrock, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + acc := headerOverrideTestAccount(tt.platform, tt.accType, nil) + require.Equal(t, tt.want, acc.IsHeaderOverrideEligible()) + }) + } + + var nilAccount *Account + require.False(t, nilAccount.IsHeaderOverrideEligible()) + require.False(t, nilAccount.IsHeaderOverrideEnabled()) + require.Nil(t, nilAccount.GetHeaderOverrides()) +} + +func TestIsHeaderOverrideEnabled(t *testing.T) { + acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }) + require.True(t, acc.IsHeaderOverrideEnabled()) + + // 未配置 / 非 bool / false 均视为未启用 + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, nil).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: "true", + }).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: false, + }).IsHeaderOverrideEnabled()) + + // 不符合平台/类型条件时即使配置了 true 也不启用 + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformGemini, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }).IsHeaderOverrideEnabled()) +} + +func TestGetHeaderOverrides(t *testing.T) { + acc := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "User-Agent": "my-agent/1.0", // 大写 key 归一化为小写 + " X-App ": "cli", // 名称去空白 + "x-empty": "", // 空 value(模板占位)跳过 + "authorization": "Bearer leaked", // 禁止覆写的头跳过 + "bad name": "value", // 非法 header 名跳过 + "x-padded": " padded ", // value 去空白 + }, + }) + overrides := acc.GetHeaderOverrides() + require.Equal(t, map[string]string{ + "user-agent": "my-agent/1.0", + "x-app": "cli", + "x-padded": "padded", + }, overrides) + + // 未启用时返回 nil + disabled := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrides: map[string]any{"user-agent": "x"}, + }) + require.Nil(t, disabled.GetHeaderOverrides()) + + // 启用但全部为空 value 时返回 nil + empty := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"user-agent": ""}, + }) + require.Nil(t, empty.GetHeaderOverrides()) + + // 未经 Normalize 落库的超长数据 / WebSocket 握手头在应用时被防御性跳过 + oversizedValue := strings.Repeat("a", maxHeaderOverrideValueLength+1) + defensive := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "x-big": oversizedValue, + "sec-websocket-key": "forged", + "content-type": "application/json", // 名单扩充前落库的数据也要被拦截 + "x-claude-code-session-id": "pinned-session", + "x-ok": "ok", + }, + }) + require.Equal(t, map[string]string{"x-ok": "ok"}, defensive.GetHeaderOverrides()) +} + +func TestApplyHeaderOverrides(t *testing.T) { + acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "user-agent": "override-agent/2.0", + "anthropic-beta": "custom-beta-1", + "x-custom": "custom-value", + }, + }) + + h := http.Header{} + // 模拟转发链路:canonical key 与 wire casing 原样 key 混合存在 + h.Set("User-Agent", "claude-cli/2.1.161 (external, cli)") + h["anthropic-beta"] = []string{"claude-code-20250219,oauth-2025-04-20"} // 非 canonical 原样 key + h.Set("Content-Type", "application/json") + + acc.ApplyHeaderOverrides(h) + + // user-agent 覆盖且只有一个值(已知头恢复 wire casing) + require.Equal(t, []string{"override-agent/2.0"}, h["User-Agent"]) + // anthropic-beta:非 canonical 旧值被清除,写入 wire casing(小写) + require.Equal(t, []string{"custom-beta-1"}, h["anthropic-beta"]) + require.Empty(t, h["Anthropic-Beta"]) + // 新增头(未知头以小写原样键写入,与转发链路 wire casing 约定一致) + require.Equal(t, []string{"custom-value"}, h["x-custom"]) + require.Equal(t, "custom-value", getHeaderRaw(h, "x-custom")) + // 未覆写的头不受影响 + require.Equal(t, "application/json", h.Get("Content-Type")) + + // 覆盖后不存在任何大小写重复 + count := 0 + for k := range h { + if k == "anthropic-beta" || k == "Anthropic-Beta" { + count++ + } + } + require.Equal(t, 1, count) +} + +func TestApplyHeaderOverridesNoOpPaths(t *testing.T) { + baseline := func() http.Header { + h := http.Header{} + h.Set("User-Agent", "orig") + return h + } + + // OAuth 账号:即使配置了覆写也不生效 + oauth := headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"}, + }) + h := baseline() + oauth.ApplyHeaderOverrides(h) + require.Equal(t, "orig", h.Get("User-Agent")) + + // 未启用开关 + off := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"}, + }) + h = baseline() + off.ApplyHeaderOverrides(h) + require.Equal(t, "orig", h.Get("User-Agent")) + + // 禁止覆写的头(authorization / x-api-key / host 等)不会被应用 + blocked := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "Authorization": "Bearer evil", + "X-Api-Key": "evil", + "Host": "evil.example.com", + "Content-Length": "0", + }, + }) + h = http.Header{} + h.Set("Authorization", "Bearer real-key") + blocked.ApplyHeaderOverrides(h) + require.Equal(t, "Bearer real-key", h.Get("Authorization")) + require.Empty(t, h.Get("X-Api-Key")) + require.Empty(t, h.Get("Host")) + + // nil header 不 panic + blocked.ApplyHeaderOverrides(nil) +} + +func TestNormalizeHeaderOverrideCredentials(t *testing.T) { + t.Run("nil credentials no-op", func(t *testing.T) { + require.NoError(t, NormalizeHeaderOverrideCredentials(nil)) + }) + + t.Run("missing keys no-op", func(t *testing.T) { + creds := map[string]any{"api_key": "sk-xxx"} + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + _, exists := creds[credKeyHeaderOverrides] + require.False(t, exists) + }) + + t.Run("normalizes names and values", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + " User-Agent ": " my-agent ", + "X-App": "", + "": "", // 完全空行被丢弃 + }, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{ + "user-agent": "my-agent", + "x-app": "", + }, creds[credKeyHeaderOverrides]) + }) + + t.Run("accepts map[string]string input", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrides: map[string]string{"X-App": "cli"}, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{"x-app": "cli"}, creds[credKeyHeaderOverrides]) + }) + + t.Run("rejects non-bool enabled", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrideEnabled: "yes", + }) + require.Error(t, err) + }) + + t.Run("rejects non-object overrides", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: []any{"user-agent"}, + }) + require.Error(t, err) + }) + + t.Run("rejects non-string value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": 123}, + }) + require.Error(t, err) + }) + + t.Run("rejects invalid header name", func(t *testing.T) { + for _, name := range []string{"bad name", "bad:name", "bad\nname", "值"} { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{name: "v"}, + }) + require.Error(t, err, "name %q should be rejected", name) + } + }) + + t.Run("rejects empty name with value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{" ": "v"}, + }) + require.Error(t, err) + }) + + t.Run("rejects blocked headers", func(t *testing.T) { + for _, name := range []string{ + "Authorization", "x-api-key", "Host", "content-length", "Transfer-Encoding", + "connection", "accept-encoding", "Sec-WebSocket-Key", "session_id", + "conversation_id", "x-codex-turn-state", "chatgpt-account-id", + "Content-Type", "Cookie", "x-goog-api-key", + "X-Claude-Code-Session-Id", "x-client-request-id", + } { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{name: "v"}, + }) + require.Error(t, err, "blocked header %q should be rejected", name) + } + }) + + t.Run("allows tab inside value", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": "a\tb"}, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{"x-app": "a\tb"}, creds[credKeyHeaderOverrides]) + }) + + t.Run("rejects invalid value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": "bad\nvalue"}, + }) + require.Error(t, err) + }) + + t.Run("rejects duplicate names case-insensitively", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{ + "User-Agent": "a", + "user-agent": "b", + }, + }) + require.Error(t, err) + }) + + t.Run("rejects too many entries", func(t *testing.T) { + entries := make(map[string]any, maxHeaderOverrideEntries+1) + for i := 0; i <= maxHeaderOverrideEntries; i++ { + entries["x-h-"+string(rune('a'+i%26))+string(rune('a'+(i/26)%26))+string(rune('a'+(i/676)%26))] = "v" + } + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: entries, + }) + require.Error(t, err) + }) + + t.Run("rejects oversized value", func(t *testing.T) { + big := make([]byte, maxHeaderOverrideValueLength+1) + for i := range big { + big[i] = 'a' + } + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": string(big)}, + }) + require.Error(t, err) + }) +} diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index dcba614c2c..5956684f98 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -39,6 +39,9 @@ type AccountRepository interface { List(ctx context.Context, params pagination.PaginationParams) ([]Account, *pagination.PaginationResult, error) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) + // ListAllWithFilters 返回符合过滤条件的全部账号(不分页),用于账号列表页 + // 计算 OpenAI 调度分数的过滤范围池。 + ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) ListActive(ctx context.Context) ([]Account, error) ListOAuthRefreshCandidates(ctx context.Context) ([]Account, error) diff --git a/backend/internal/service/account_service_delete_test.go b/backend/internal/service/account_service_delete_test.go index a304356c09..ee6163239e 100644 --- a/backend/internal/service/account_service_delete_test.go +++ b/backend/internal/service/account_service_delete_test.go @@ -79,6 +79,10 @@ func (s *accountRepoStub) List(ctx context.Context, params pagination.Pagination panic("unexpected List call") } +func (s *accountRepoStub) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStub) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected ListWithFilters call") } diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 80a862b971..7bf02afd51 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -295,6 +295,9 @@ func (s *AccountTestService) testClaudeAccountConnection(c *gin.Context, account setAnthropicAPIKeyAuthHeader(req.Header, account, authToken) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + // Get proxy URL proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -600,9 +603,19 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account if isOAuth { req.Host = "chatgpt.com" req.Header.Set("accept", "text/event-stream") + req.Header.Set("OpenAI-Beta", "responses=experimental") + req.Header.Set("Originator", "codex_cli_rs") + if customUA := strings.TrimSpace(credentialAccount.GetOpenAIUserAgent()); customUA != "" { + req.Header.Set("User-Agent", customUA) + } else { + req.Header.Set("User-Agent", codexCLIUserAgent) + } setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + credentialAccount.ApplyHeaderOverrides(req.Header) + // Get proxy URL proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -756,6 +769,9 @@ func (s *AccountTestService) testOpenAIChatCompletionsConnection( req.Header.Set("Accept", "text/event-stream") req.Header.Set("Authorization", "Bearer "+authToken) + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -848,6 +864,9 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account setOpenAIChatGPTAccountHeaders(req.Header, account) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -1599,6 +1618,9 @@ func (s *AccountTestService) testOpenAIImageAPIKey(c *gin.Context, ctx context.C req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+authToken) + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 3c50baaec2..d9f6200359 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -184,6 +184,7 @@ type UsageInfo struct { FiveHour *UsageProgress `json:"five_hour"` // 5小时窗口 SevenDay *UsageProgress `json:"seven_day,omitempty"` // 7天窗口 SevenDaySonnet *UsageProgress `json:"seven_day_sonnet,omitempty"` // 7天Sonnet窗口 + SevenDayFable *UsageProgress `json:"seven_day_fable,omitempty"` // 7天Fable窗口(响应头 7d_oi) GeminiSharedDaily *UsageProgress `json:"gemini_shared_daily,omitempty"` // Gemini shared pool RPD (Google One / Code Assist) GeminiProDaily *UsageProgress `json:"gemini_pro_daily,omitempty"` // Gemini Pro 日配额 GeminiFlashDaily *UsageProgress `json:"gemini_flash_daily,omitempty"` // Gemini Flash 日配额 @@ -236,6 +237,12 @@ type UsageInfo struct { Error string `json:"error,omitempty"` } +// ClaudeUsageWindow Anthropic /api/oauth/usage 返回的单个用量窗口 +type ClaudeUsageWindow struct { + Utilization float64 `json:"utilization"` + ResetsAt string `json:"resets_at"` +} + // ClaudeUsageResponse Anthropic API返回的usage结构 type ClaudeUsageResponse struct { FiveHour struct { @@ -250,6 +257,10 @@ type ClaudeUsageResponse struct { Utilization float64 `json:"utilization"` ResetsAt string `json:"resets_at"` } `json:"seven_day_sonnet"` + // Fable 专属 7d 窗口(对应响应头 7d_oi,claim 名为 seven_day_overage_included, + // 见 anthropic-ratelimit-unified-representative-claim 头)。上游 usage API + // 若不下发该字段,GetUsage 会用被动采样数据回填。 + SevenDayOverageIncluded ClaudeUsageWindow `json:"seven_day_overage_included"` } // ClaudeUsageFetchOptions 包含获取 Claude 用量数据所需的所有选项 @@ -429,6 +440,12 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for // 5. 将主动查询结果同步到被动缓存,下次 passive 加载即为最新值 s.syncActiveToPassive(ctx, account.ID, usage) + // 6. 上游 usage API 目前不一定下发 Fable 7d 窗口;缺失时回填被动采样 + // (7d_oi 响应头)的数据,避免主动查询后 7d F 进度条丢失。 + if usage.SevenDayFable == nil { + usage.SevenDayFable = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset") + } + s.tryClearRecoverableAccountError(ctx, account) return usage, nil } @@ -471,25 +488,10 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int } // 构建 7d 窗口(从被动采样数据) - util7d := parseExtraFloat64(account.Extra["passive_usage_7d_utilization"]) - reset7dRaw := parseExtraFloat64(account.Extra["passive_usage_7d_reset"]) - if util7d > 0 || reset7dRaw > 0 { - var resetAt *time.Time - var remaining int - if reset7dRaw > 0 { - t := time.Unix(int64(reset7dRaw), 0) - resetAt = &t - remaining = int(time.Until(t).Seconds()) - if remaining < 0 { - remaining = 0 - } - } - info.SevenDay = &UsageProgress{ - Utilization: util7d * 100, - ResetsAt: resetAt, - RemainingSeconds: remaining, - } - } + info.SevenDay = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_utilization", "passive_usage_7d_reset") + + // 构建 7d Fable 窗口(从被动采样的 7d_oi 响应头数据) + info.SevenDayFable = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset") // 添加窗口统计 s.addWindowStats(ctx, account, info) @@ -497,6 +499,31 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int return info, nil } +// buildPassiveUsageWindow 从 Extra 中的被动采样数据(utilization 为 0-1 小数、reset 为 Unix 秒) +// 构建用量窗口,无数据时返回 nil。 +func buildPassiveUsageWindow(extra map[string]any, utilKey, resetKey string) *UsageProgress { + util := parseExtraFloat64(extra[utilKey]) + resetRaw := parseExtraFloat64(extra[resetKey]) + if util <= 0 && resetRaw <= 0 { + return nil + } + var resetAt *time.Time + var remaining int + if resetRaw > 0 { + t := time.Unix(int64(resetRaw), 0) + resetAt = &t + remaining = int(time.Until(t).Seconds()) + if remaining < 0 { + remaining = 0 + } + } + return &UsageProgress{ + Utilization: util * 100, + ResetsAt: resetAt, + RemainingSeconds: remaining, + } +} + // syncActiveToPassive 将主动查询的最新数据回写到 Extra 被动缓存, // 这样下次被动加载时能看到最新值。 func (s *AccountUsageService) syncActiveToPassive(ctx context.Context, accountID int64, usage *UsageInfo) { @@ -511,6 +538,12 @@ func (s *AccountUsageService) syncActiveToPassive(ctx context.Context, accountID extraUpdates["passive_usage_7d_reset"] = usage.SevenDay.ResetsAt.Unix() } } + if usage.SevenDayFable != nil { + extraUpdates["passive_usage_7d_oi_utilization"] = usage.SevenDayFable.Utilization / 100 + if usage.SevenDayFable.ResetsAt != nil { + extraUpdates["passive_usage_7d_oi_reset"] = usage.SevenDayFable.ResetsAt.Unix() + } + } if len(extraUpdates) > 0 { extraUpdates["passive_usage_sampled_at"] = time.Now().UTC().Format(time.RFC3339) @@ -1010,8 +1043,8 @@ func enrichUsageWithAccountError(info *UsageInfo, account *Account) { // 使用独立缓存(1 分钟),与 API 缓存分离 func (s *AccountUsageService) addWindowStats(ctx context.Context, account *Account, usage *UsageInfo) { // 修复:即使 FiveHour 为 nil,也要尝试获取统计数据 - // 因为 SevenDay/SevenDaySonnet 可能需要 - if usage.FiveHour == nil && usage.SevenDay == nil && usage.SevenDaySonnet == nil { + // 因为 SevenDay/SevenDaySonnet/SevenDayFable 可能需要 + if usage.FiveHour == nil && usage.SevenDay == nil && usage.SevenDaySonnet == nil && usage.SevenDayFable == nil { return } @@ -1347,6 +1380,22 @@ func (s *AccountUsageService) buildUsageInfo(resp *ClaudeUsageResponse, updatedA } } + // 7天Fable窗口(响应头 7d_oi 对应的窗口) + if fable := resp.SevenDayOverageIncluded; fable.ResetsAt != "" { + if fableReset, err := parseTime(fable.ResetsAt); err == nil { + info.SevenDayFable = &UsageProgress{ + Utilization: fable.Utilization, + ResetsAt: &fableReset, + RemainingSeconds: int(time.Until(fableReset).Seconds()), + } + } else { + log.Printf("Failed to parse SevenDayFable.ResetsAt: %s, error: %v", fable.ResetsAt, err) + info.SevenDayFable = &UsageProgress{ + Utilization: fable.Utilization, + } + } + } + return info } diff --git a/backend/internal/service/account_usage_service_fable_test.go b/backend/internal/service/account_usage_service_fable_test.go new file mode 100644 index 0000000000..60f2e9c79f --- /dev/null +++ b/backend/internal/service/account_usage_service_fable_test.go @@ -0,0 +1,119 @@ +package service + +import ( + "encoding/json" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestClaudeUsageResponse_FableWindowDecoding(t *testing.T) { + t.Run("seven_day_overage_included", func(t *testing.T) { + raw := `{ + "five_hour": {"utilization": 12.0, "resets_at": "2026-07-03T10:00:00Z"}, + "seven_day": {"utilization": 34.0, "resets_at": "2026-07-08T00:00:00Z"}, + "seven_day_overage_included": {"utilization": 56.0, "resets_at": "2026-07-08T03:00:00Z"} +}` + var resp ClaudeUsageResponse + require.NoError(t, json.Unmarshal([]byte(raw), &resp)) + require.Equal(t, 56.0, resp.SevenDayOverageIncluded.Utilization) + require.Equal(t, "2026-07-08T03:00:00Z", resp.SevenDayOverageIncluded.ResetsAt) + }) + + t.Run("absent", func(t *testing.T) { + raw := `{"five_hour": {"utilization": 12.0, "resets_at": "2026-07-03T10:00:00Z"}}` + var resp ClaudeUsageResponse + require.NoError(t, json.Unmarshal([]byte(raw), &resp)) + require.Zero(t, resp.SevenDayOverageIncluded.Utilization) + require.Empty(t, resp.SevenDayOverageIncluded.ResetsAt) + }) +} + +func TestBuildUsageInfo_SevenDayFable(t *testing.T) { + svc := &AccountUsageService{} + now := time.Now() + + resetAt := now.Add(72 * time.Hour).UTC().Truncate(time.Second) + var resp ClaudeUsageResponse + resp.FiveHour.Utilization = 10 + resp.SevenDayOverageIncluded = ClaudeUsageWindow{ + Utilization: 88, + ResetsAt: resetAt.Format(time.RFC3339), + } + + info := svc.buildUsageInfo(&resp, &now) + require.NotNil(t, info.SevenDayFable) + require.Equal(t, 88.0, info.SevenDayFable.Utilization) + require.NotNil(t, info.SevenDayFable.ResetsAt) + require.True(t, info.SevenDayFable.ResetsAt.Equal(resetAt)) + require.Greater(t, info.SevenDayFable.RemainingSeconds, 0) + + // 无 Fable 数据时不应创建窗口 + var empty ClaudeUsageResponse + empty.FiveHour.Utilization = 10 + info = svc.buildUsageInfo(&empty, &now) + require.Nil(t, info.SevenDayFable) +} + +func TestBuildPassiveUsageWindow(t *testing.T) { + future := time.Now().Add(48 * time.Hour).Unix() + + t.Run("utilization and reset", func(t *testing.T) { + window := buildPassiveUsageWindow(map[string]any{ + "passive_usage_7d_oi_utilization": 0.87, + "passive_usage_7d_oi_reset": float64(future), + }, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset") + require.NotNil(t, window) + require.InDelta(t, 87.0, window.Utilization, 1e-9) + require.NotNil(t, window.ResetsAt) + require.Equal(t, future, window.ResetsAt.Unix()) + require.Greater(t, window.RemainingSeconds, 0) + }) + + t.Run("no data returns nil", func(t *testing.T) { + require.Nil(t, buildPassiveUsageWindow(nil, "u", "r")) + require.Nil(t, buildPassiveUsageWindow(map[string]any{}, "u", "r")) + }) + + t.Run("expired reset clamps remaining to zero", func(t *testing.T) { + past := time.Now().Add(-time.Hour).Unix() + window := buildPassiveUsageWindow(map[string]any{ + "u": 0.5, + "r": float64(past), + }, "u", "r") + require.NotNil(t, window) + require.Equal(t, 0, window.RemainingSeconds) + }) + + t.Run("utilization only", func(t *testing.T) { + window := buildPassiveUsageWindow(map[string]any{"u": 0.25}, "u", "r") + require.NotNil(t, window) + require.InDelta(t, 25.0, window.Utilization, 1e-9) + require.Nil(t, window.ResetsAt) + }) +} + +func TestSyncActiveToPassive_WritesFableExtras(t *testing.T) { + repo := &accountUsageCodexProbeRepo{updateExtraCh: make(chan map[string]any, 1)} + svc := &AccountUsageService{accountRepo: repo} + + resetAt := time.Now().Add(72 * time.Hour).Truncate(time.Second) + usage := &UsageInfo{ + SevenDayFable: &UsageProgress{ + Utilization: 87, + ResetsAt: &resetAt, + }, + } + + svc.syncActiveToPassive(t.Context(), 1, usage) + + select { + case updates := <-repo.updateExtraCh: + require.InDelta(t, 0.87, updates["passive_usage_7d_oi_utilization"], 1e-9) + require.Equal(t, resetAt.Unix(), updates["passive_usage_7d_oi_reset"]) + require.Contains(t, updates, "passive_usage_sampled_at") + default: + t.Fatal("expected UpdateExtra to be called with fable extras") + } +} diff --git a/backend/internal/service/account_wildcard_test.go b/backend/internal/service/account_wildcard_test.go index d903b940a5..6ce804bb16 100644 --- a/backend/internal/service/account_wildcard_test.go +++ b/backend/internal/service/account_wildcard_test.go @@ -4,6 +4,8 @@ package service import ( "testing" + + "github.com/Wei-Shaw/sub2api/internal/domain" ) func TestMatchWildcard(t *testing.T) { @@ -320,6 +322,86 @@ func TestAccountGetMappedModel(t *testing.T) { } } +func TestAccountGetModelMapping_AntigravityNormalizesGemini31ProAliases(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformAntigravity, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-high": "gemini-3.1-pro-high", + "gemini-3.1-pro-preview": "gemini-3.1-pro-high", + }, + }, + } + + mapping := account.GetModelMapping() + + if got := mapping["gemini-3.1-pro"]; got != domain.AntigravityGemini31ProAgentModel { + t.Fatalf("expected gemini-3.1-pro to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got) + } + if got := mapping["gemini-3.1-pro-high"]; got != domain.AntigravityGemini31ProAgentModel { + t.Fatalf("expected gemini-3.1-pro-high to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got) + } + if got := mapping["gemini-3.1-pro-preview"]; got != domain.AntigravityGemini31ProAgentModel { + t.Fatalf("expected gemini-3.1-pro-preview to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got) + } +} + +func TestAccountGetModelMapping_AntigravityPreservesGemini31ProOverrides(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformAntigravity, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel, + "gemini-3.1-pro-high": "custom-high", + "gemini-3.1-pro-preview": "custom-preview", + }, + }, + } + + mapping := account.GetModelMapping() + + if got := mapping["gemini-3.1-pro-high"]; got != "custom-high" { + t.Fatalf("expected gemini-3.1-pro-high override to be preserved, got %q", got) + } + if got := mapping["gemini-3.1-pro-preview"]; got != "custom-preview" { + t.Fatalf("expected gemini-3.1-pro-preview override to be preserved, got %q", got) + } + if got := mapping["gemini-3.1-pro"]; got != domain.AntigravityGemini31ProAgentModel { + t.Fatalf("expected gemini-3.1-pro alias to default to %q, got %q", domain.AntigravityGemini31ProAgentModel, got) + } +} + +func TestAccountGetModelMapping_AntigravityGemini31ProAliasesRespectWildcard(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformAntigravity, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel, + "gemini-3.1-*": "custom-wildcard", + }, + }, + } + + mapping := account.GetModelMapping() + + if got := mapping["gemini-3.1-pro"]; got != "" { + t.Fatalf("expected gemini-3.1-pro exact alias to stay unset when wildcard exists, got %q", got) + } + if got := mapping["gemini-3.1-pro-high"]; got != "" { + t.Fatalf("expected gemini-3.1-pro-high exact alias to stay unset when wildcard exists, got %q", got) + } + if got := mapping["gemini-3.1-pro-preview"]; got != "" { + t.Fatalf("expected gemini-3.1-pro-preview exact alias to stay unset when wildcard exists, got %q", got) + } +} + func TestAccountResolveMappedModel(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/admin_account_concurrency_test.go b/backend/internal/service/admin_account_concurrency_test.go index 3544f80e24..da57b5a64f 100644 --- a/backend/internal/service/admin_account_concurrency_test.go +++ b/backend/internal/service/admin_account_concurrency_test.go @@ -5,23 +5,16 @@ package service import ( "testing" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/stretchr/testify/require" ) -func TestNormalizeAccountConcurrencyCapsGrokOAuthUnlessUnsafe(t *testing.T) { - t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "") - +func TestNormalizeAccountConcurrencyDefaultsInvalidGrokOAuthToOne(t *testing.T) { require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0)) require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, -5)) - require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50)) +} + +func TestNormalizeAccountConcurrencyPreservesExplicitValues(t *testing.T) { + require.Equal(t, 50, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50)) require.Equal(t, 2, normalizeAccountConcurrency(PlatformOpenAI, AccountTypeOAuth, 2)) require.Equal(t, 2, normalizeAccountConcurrency(PlatformGrok, AccountTypeAPIKey, 2)) } - -func TestNormalizeAccountConcurrencyAllowsGrokOAuthUnsafeOverride(t *testing.T) { - t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "true") - - require.Equal(t, 50, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50)) - require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0)) -} diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index bacd134db4..f1de60eb47 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -78,6 +78,12 @@ type AdminService interface { // Account management ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error) + // ListAccountsForSchedulerScoreFilter 返回符合过滤条件的全部账号(不分页), + // 作为账号列表页计算 OpenAI 调度分数的过滤范围池。 + ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) + // ListOpenAISchedulableAccountsForSchedulerScore 返回指定分组(nil 为未分组)内 + // 可调度的 OpenAI 账号,用于按组计算调度分数。 + ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) GetAccount(ctx context.Context, id int64) (*Account, error) GetAccountsByIDs(ctx context.Context, ids []int64) ([]*Account, error) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) @@ -2618,6 +2624,23 @@ func (s *adminServiceImpl) ListAccounts(ctx context.Context, page, pageSize int, return accounts, result.Total, nil } +func (s *adminServiceImpl) ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + if s == nil || s.accountRepo == nil { + return nil, nil + } + return s.accountRepo.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode) +} + +func (s *adminServiceImpl) ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) { + if s == nil || s.accountRepo == nil { + return nil, nil + } + if groupID != nil { + return s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, PlatformOpenAI) + } + return s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, PlatformOpenAI) +} + func (s *adminServiceImpl) GetAccount(ctx context.Context, id int64) (*Account, error) { return s.accountRepo.GetByID(ctx, id) } @@ -2640,9 +2663,6 @@ func normalizeAccountConcurrency(platform, accountType string, concurrency int) if concurrency <= 0 { return 1 } - if concurrency > 1 && !xai.AllowUnsafeHighConcurrency() { - return 1 - } } return concurrency } @@ -2671,6 +2691,11 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou } } + // 校验并规范化请求头覆写配置(header 名小写化、格式检查) + if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { + return nil, err + } + account := &Account{ Name: input.Name, Notes: normalizeAccountNotes(input.Notes), @@ -2801,6 +2826,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // 敏感子键采用"incoming 没提供就保留"的合并语义:前端响应已脱敏, // 全对象 PUT 编辑时不会再带回 token,避免覆盖时清空已有凭证。 account.Credentials = MergePreservingSensitiveCreds(account.Credentials, input.Credentials) + // 校验并规范化请求头覆写配置(header 名小写化、格式检查) + if err := NormalizeHeaderOverrideCredentials(account.Credentials); err != nil { + return nil, err + } } // Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。 // 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。 @@ -3019,6 +3048,11 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp } } + // 校验并规范化请求头覆写配置(批量路径为 JSONB 顶层 key 合并,直接校验增量即可) + if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { + return nil, err + } + // Prepare bulk updates for columns and JSONB fields. repoUpdates := AccountBulkUpdate{ Credentials: input.Credentials, diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index df415295b1..2f44b1741d 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -88,6 +88,10 @@ func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID in return nil, nil } +func (s *accountRepoStubForBulkUpdate) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForBulkUpdate) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listCalled = true s.lastListParams = params diff --git a/backend/internal/service/admin_service_search_test.go b/backend/internal/service/admin_service_search_test.go index 595e99e344..76acd1b5e2 100644 --- a/backend/internal/service/admin_service_search_test.go +++ b/backend/internal/service/admin_service_search_test.go @@ -25,6 +25,10 @@ type accountRepoStubForAdminList struct { listWithFiltersErr error } +func (s *accountRepoStubForAdminList) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForAdminList) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listWithFiltersCalls++ s.listWithFiltersParams = params diff --git a/backend/internal/service/antigravity_token_refresher.go b/backend/internal/service/antigravity_token_refresher.go index 7ce0ccf0fe..f40df7d95f 100644 --- a/backend/internal/service/antigravity_token_refresher.go +++ b/backend/internal/service/antigravity_token_refresher.go @@ -12,6 +12,10 @@ const ( // antigravityRefreshWindow Antigravity token 提前刷新窗口:15分钟 // Google OAuth token 有效期55分钟,提前15分钟刷新 antigravityRefreshWindow = 15 * time.Minute + + antigravityForceTokenRefreshExtraKey = "antigravity_force_token_refresh" + antigravityForceTokenRefreshReasonExtraKey = "antigravity_force_token_refresh_reason" + antigravityForceTokenRefreshAtExtraKey = "antigravity_force_token_refresh_at" ) // AntigravityTokenRefresher 实现 TokenRefresher 接口 @@ -41,6 +45,9 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati if !r.CanRefresh(account) { return false } + if accountNeedsAntigravityForceTokenRefresh(account) { + return true + } expiresAt := account.GetCredentialAsTime("expires_at") if expiresAt == nil { return false @@ -54,6 +61,29 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati return needsRefresh } +func accountNeedsAntigravityForceTokenRefresh(account *Account) bool { + return account != nil && + account.Platform == PlatformAntigravity && + account.Type == AccountTypeOAuth && + account.getExtraBool(antigravityForceTokenRefreshExtraKey) +} + +func antigravityForceTokenRefreshExtra(reason string) map[string]any { + return map[string]any{ + antigravityForceTokenRefreshExtraKey: true, + antigravityForceTokenRefreshReasonExtraKey: reason, + antigravityForceTokenRefreshAtExtraKey: time.Now().UTC().Format(time.RFC3339), + } +} + +func clearAntigravityForceTokenRefreshExtra() map[string]any { + return map[string]any{ + antigravityForceTokenRefreshExtraKey: false, + antigravityForceTokenRefreshReasonExtraKey: "", + antigravityForceTokenRefreshAtExtraKey: "", + } +} + // Refresh 执行 token 刷新 func (r *AntigravityTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) { tokenInfo, err := r.antigravityOAuthService.RefreshAccountToken(ctx, account) diff --git a/backend/internal/service/api_key.go b/backend/internal/service/api_key.go index ec20b0a9bf..dfc3ec1c5a 100644 --- a/backend/internal/service/api_key.go +++ b/backend/internal/service/api_key.go @@ -44,6 +44,7 @@ type APIKey struct { UpdatedAt time.Time User *User Group *Group + CurrentConcurrency int // Quota fields Quota float64 // Quota limit in USD (0 = unlimited) diff --git a/backend/internal/service/api_key_service.go b/backend/internal/service/api_key_service.go index de9b908dc7..8903be65ee 100644 --- a/backend/internal/service/api_key_service.go +++ b/backend/internal/service/api_key_service.go @@ -203,6 +203,7 @@ type APIKeyService struct { userGroupRateRepo UserGroupRateRepository cache APIKeyCache rateLimitCacheInvalid RateLimitCacheInvalidator // optional: invalidate Redis rate limit cache + concurrencyService *ConcurrencyService cfg *config.Config authCacheL1 *ristretto.Cache authCfg apiKeyAuthCacheConfig @@ -240,6 +241,10 @@ func (s *APIKeyService) SetRateLimitCacheInvalidator(inv RateLimitCacheInvalidat s.rateLimitCacheInvalid = inv } +func (s *APIKeyService) SetConcurrencyService(concurrencyService *ConcurrencyService) { + s.concurrencyService = concurrencyService +} + func (s *APIKeyService) compileAPIKeyIPRules(apiKey *APIKey) { if apiKey == nil { return @@ -436,9 +441,40 @@ func (s *APIKeyService) List(ctx context.Context, userID int64, params paginatio if err != nil { return nil, nil, fmt.Errorf("list api keys: %w", err) } + s.fillCurrentConcurrency(ctx, keys) return keys, pagination, nil } +func (s *APIKeyService) fillCurrentConcurrency(ctx context.Context, keys []APIKey) { + if s == nil || s.concurrencyService == nil || len(keys) == 0 { + return + } + ids := make([]int64, 0, len(keys)) + for i := range keys { + if keys[i].ID > 0 { + ids = append(ids, keys[i].ID) + } + } + counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, ids) + if err != nil { + return + } + for i := range keys { + keys[i].CurrentConcurrency = counts[keys[i].ID] + } +} + +func (s *APIKeyService) currentConcurrencyForAPIKey(ctx context.Context, apiKeyID int64) int { + if s == nil || s.concurrencyService == nil || apiKeyID <= 0 { + return 0 + } + counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, []int64{apiKeyID}) + if err != nil { + return 0 + } + return counts[apiKeyID] +} + func (s *APIKeyService) VerifyOwnership(ctx context.Context, userID int64, apiKeyIDs []int64) ([]int64, error) { if len(apiKeyIDs) == 0 { return []int64{}, nil @@ -458,6 +494,9 @@ func (s *APIKeyService) GetByID(ctx context.Context, id int64) (*APIKey, error) return nil, fmt.Errorf("get api key: %w", err) } s.compileAPIKeyIPRules(apiKey) + if apiKey != nil { + apiKey.CurrentConcurrency = s.currentConcurrencyForAPIKey(ctx, apiKey.ID) + } return apiKey, nil } diff --git a/backend/internal/service/api_key_service_delete_test.go b/backend/internal/service/api_key_service_delete_test.go index 8664c03bd7..25ad1edb15 100644 --- a/backend/internal/service/api_key_service_delete_test.go +++ b/backend/internal/service/api_key_service_delete_test.go @@ -300,6 +300,40 @@ func TestApiKeyService_Delete_NotFound(t *testing.T) { require.Empty(t, cache.deleteAuthKeys) } +func TestAPIKeyService_List_FillsCurrentConcurrency(t *testing.T) { + repo := &apiKeyRepoStub{ + allowListByUserID: true, + listByUserIDKeys: []APIKey{ + {ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"}, + {ID: 11, UserID: 7, Key: "sk-11", Name: "key-11"}, + }, + } + concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{ + apiKeyConcurrency: map[int64]int{10: 2, 11: 0}, + }) + svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency} + + keys, _, err := svc.List(context.Background(), 7, pagination.PaginationParams{Page: 1, PageSize: 20}, APIKeyListFilters{}) + require.NoError(t, err) + require.Len(t, keys, 2) + require.Equal(t, 2, keys[0].CurrentConcurrency) + require.Equal(t, 0, keys[1].CurrentConcurrency) +} + +func TestAPIKeyService_GetByID_FillsCurrentConcurrency(t *testing.T) { + repo := &apiKeyRepoStub{ + apiKey: &APIKey{ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"}, + } + concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{ + apiKeyConcurrency: map[int64]int{10: 4}, + }) + svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency} + + key, err := svc.GetByID(context.Background(), 10) + require.NoError(t, err) + require.Equal(t, 4, key.CurrentConcurrency) +} + // TestApiKeyService_Delete_DeleteFails 测试删除操作失败时的错误处理。 // 预期行为: // - GetKeyAndOwnerID 返回正确的所有者 ID diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index a781936598..dc54a1b1f3 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -280,6 +280,11 @@ func (s *BillingService) initFallbackPricing() { s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"] s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"] + // GPT-5.6(sol / terra / luna)暂无独立定价,回退到 GPT-5.4。 + s.fallbackPrices["gpt-5.6-sol"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.6-terra"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.6-luna"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{ InputPricePerToken: 7.5e-7, OutputPricePerToken: 4.5e-6, @@ -667,6 +672,12 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { // OpenAI(GPT-5 / Codex 族):仅匹配已知型号,避免未知 OpenAI 型号误计价。 if normalized := normalizeKnownOpenAICodexModel(modelLower); normalized != "" { switch normalized { + case "gpt-5.6-sol": + return s.fallbackPrices["gpt-5.6-sol"] + case "gpt-5.6-terra": + return s.fallbackPrices["gpt-5.6-terra"] + case "gpt-5.6-luna": + return s.fallbackPrices["gpt-5.6-luna"] case "gpt-5.5-pro": return s.fallbackPrices["gpt-5.5-pro"] case "gpt-5.5": @@ -1060,7 +1071,8 @@ func isOpenAIGPT54Model(model string) bool { // normalizeCodexModel 的默认兜底把非 OpenAI 模型(claude-*、gemini-*、gpt-4o) // 误识别为 gpt-5.4。 normalized := normalizeKnownOpenAICodexModel(model) - return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" + return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" || + normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" } // CalculateCostWithConfig 使用配置中的默认倍率计算费用 diff --git a/backend/internal/service/codex_image_generation_bridge.go b/backend/internal/service/codex_image_generation_bridge.go index c7a894a792..68989d67dc 100644 --- a/backend/internal/service/codex_image_generation_bridge.go +++ b/backend/internal/service/codex_image_generation_bridge.go @@ -4,6 +4,13 @@ import "strings" const featureKeyCodexImageGenerationBridge = "codex_image_generation_bridge" +const ( + featureKeyCodexImageGenerationExplicitToolPolicy = "codex_image_generation_explicit_tool_policy" + + codexImageGenerationExplicitToolPolicyAllow = "allow" + codexImageGenerationExplicitToolPolicyStrip = "strip" +) + func boolOverridePtr(v bool) *bool { return &v } @@ -20,6 +27,27 @@ func boolOverrideFromMap(values map[string]any, keys ...string) *bool { return nil } +func stringOverrideFromMap(values map[string]any, keys ...string) (string, bool) { + if values == nil { + return "", false + } + for _, key := range keys { + if v, ok := values[key].(string); ok { + return v, true + } + } + return "", false +} + +func normalizeCodexImageGenerationExplicitToolPolicy(value string) string { + switch strings.ToLower(strings.TrimSpace(value)) { + case codexImageGenerationExplicitToolPolicyStrip, "remove", "drop": + return codexImageGenerationExplicitToolPolicyStrip + default: + return codexImageGenerationExplicitToolPolicyAllow + } +} + func platformBoolOverride(values map[string]any, key string, platform string) *bool { if values == nil { return nil @@ -62,3 +90,20 @@ func (a *Account) CodexImageGenerationBridgeOverride() *bool { openaiConfig, _ := a.Extra[PlatformOpenAI].(map[string]any) return boolOverrideFromMap(openaiConfig, featureKeyCodexImageGenerationBridge, "codex_image_generation_bridge_enabled") } + +// CodexImageGenerationExplicitToolPolicy returns the account-level policy for +// client-provided Codex /responses image_generation tools. Unknown or unset +// values default to allow to preserve existing behavior. +func (a *Account) CodexImageGenerationExplicitToolPolicy() string { + if a == nil || a.Platform != PlatformOpenAI || a.Extra == nil { + return codexImageGenerationExplicitToolPolicyAllow + } + if policy, ok := stringOverrideFromMap(a.Extra, featureKeyCodexImageGenerationExplicitToolPolicy); ok { + return normalizeCodexImageGenerationExplicitToolPolicy(policy) + } + openaiConfig, _ := a.Extra[PlatformOpenAI].(map[string]any) + if policy, ok := stringOverrideFromMap(openaiConfig, featureKeyCodexImageGenerationExplicitToolPolicy); ok { + return normalizeCodexImageGenerationExplicitToolPolicy(policy) + } + return codexImageGenerationExplicitToolPolicyAllow +} diff --git a/backend/internal/service/concurrency_service.go b/backend/internal/service/concurrency_service.go index 10f47f5b74..f2f2aade89 100644 --- a/backend/internal/service/concurrency_service.go +++ b/backend/internal/service/concurrency_service.go @@ -53,6 +53,12 @@ type ConcurrencyCache interface { CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error } +type APIKeyConcurrencyCache interface { + TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) +} + var ( requestIDPrefix = initRequestIDPrefix() requestIDCounter atomic.Uint64 @@ -90,6 +96,8 @@ const ( defaultAccountLoadBatchCacheTTL = 200 * time.Millisecond accountLoadBatchFetchTimeout = 3 * time.Second maxAccountLoadBatchCacheEntries = 256 + apiKeyConcurrencyFetchTimeout = 3 * time.Second + apiKeySlotTrackTimeout = 2 * time.Second ) // ConcurrencyService 管理账号和用户的并发限制。 @@ -238,6 +246,77 @@ func (s *ConcurrencyService) AcquireUserSlot(ctx context.Context, userID int64, }, nil } +// TrackAPIKeySlot records one active request slot for an API key without +// applying key-level concurrency limits. It is fail-open: Redis errors are +// logged and return a no-op release function. +func (s *ConcurrencyService) TrackAPIKeySlot(ctx context.Context, apiKeyID int64) func() { + if s == nil || s.cache == nil || apiKeyID <= 0 { + return func() {} + } + cache, ok := s.cache.(APIKeyConcurrencyCache) + if !ok { + return func() {} + } + + requestID := generateRequestID() + baseCtx := context.Background() + if ctx != nil { + baseCtx = context.WithoutCancel(ctx) + } + trackCtx, cancel := context.WithTimeout(baseCtx, apiKeySlotTrackTimeout) + err := cache.TrackAPIKeySlot(trackCtx, apiKeyID, requestID) + cancel() + if err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: failed to track api key slot for %d (req=%s): %v", apiKeyID, requestID, err) + return func() {} + } + + return func() { + bgCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := cache.ReleaseAPIKeySlot(bgCtx, apiKeyID, requestID); err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: failed to release api key slot for %d (req=%s): %v", apiKeyID, requestID, err) + } + } +} + +// GetAPIKeyConcurrencyBatch gets real-time active request counts for API keys. +// Stats are best-effort: missing Redis support or Redis errors return zeroes. +func (s *ConcurrencyService) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + result := zeroAPIKeyConcurrencyMap(apiKeyIDs) + if len(apiKeyIDs) == 0 { + return result, nil + } + if s == nil || s.cache == nil { + return result, nil + } + cache, ok := s.cache.(APIKeyConcurrencyCache) + if !ok { + return result, nil + } + + redisCtx, cancel := context.WithTimeout(context.Background(), apiKeyConcurrencyFetchTimeout) + defer cancel() + + counts, err := cache.GetAPIKeyConcurrencyBatch(redisCtx, apiKeyIDs) + if err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: get api key concurrency batch failed: %v", err) + return result, nil + } + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = counts[apiKeyID] + } + return result, nil +} + +func zeroAPIKeyConcurrencyMap(apiKeyIDs []int64) map[int64]int { + result := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = 0 + } + return result +} + // ============================================ // Wait Queue Count Methods // ============================================ diff --git a/backend/internal/service/concurrency_service_test.go b/backend/internal/service/concurrency_service_test.go index bacad0245e..3f358bbe6a 100644 --- a/backend/internal/service/concurrency_service_test.go +++ b/backend/internal/service/concurrency_service_test.go @@ -16,25 +16,33 @@ import ( // stubConcurrencyCacheForTest 用于并发服务单元测试的缓存桩 type stubConcurrencyCacheForTest struct { - acquireResult bool - acquireErr error - releaseErr error - concurrency int - concurrencyErr error - waitAllowed bool - waitErr error - waitCount int - waitCountErr error - loadBatch map[int64]*AccountLoadInfo - loadBatchErr error - usersLoadBatch map[int64]*UserLoadInfo - usersLoadErr error - cleanupErr error + acquireResult bool + acquireErr error + releaseErr error + concurrency int + concurrencyErr error + waitAllowed bool + waitErr error + waitCount int + waitCountErr error + loadBatch map[int64]*AccountLoadInfo + loadBatchErr error + usersLoadBatch map[int64]*UserLoadInfo + usersLoadErr error + cleanupErr error + apiKeyTrackErr error + apiKeyReleaseErr error + apiKeyConcurrency map[int64]int + apiKeyConcurrencyErr error // 记录调用 - releasedAccountIDs []int64 - releasedRequestIDs []string - loadBatchCalls atomic.Int64 + releasedAccountIDs []int64 + releasedRequestIDs []string + loadBatchCalls atomic.Int64 + trackedAPIKeyIDs []int64 + trackedAPIKeyRequestIDs []string + releasedAPIKeyIDs []int64 + releasedAPIKeyRequestIDs []string } var _ ConcurrencyCache = (*stubConcurrencyCacheForTest)(nil) @@ -78,6 +86,26 @@ func (c *stubConcurrencyCacheForTest) ReleaseUserSlot(_ context.Context, _ int64 func (c *stubConcurrencyCacheForTest) GetUserConcurrency(_ context.Context, _ int64) (int, error) { return c.concurrency, c.concurrencyErr } +func (c *stubConcurrencyCacheForTest) TrackAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error { + c.trackedAPIKeyIDs = append(c.trackedAPIKeyIDs, apiKeyID) + c.trackedAPIKeyRequestIDs = append(c.trackedAPIKeyRequestIDs, requestID) + return c.apiKeyTrackErr +} +func (c *stubConcurrencyCacheForTest) ReleaseAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error { + c.releasedAPIKeyIDs = append(c.releasedAPIKeyIDs, apiKeyID) + c.releasedAPIKeyRequestIDs = append(c.releasedAPIKeyRequestIDs, requestID) + return c.apiKeyReleaseErr +} +func (c *stubConcurrencyCacheForTest) GetAPIKeyConcurrencyBatch(_ context.Context, apiKeyIDs []int64) (map[int64]int, error) { + if c.apiKeyConcurrencyErr != nil { + return nil, c.apiKeyConcurrencyErr + } + result := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = c.apiKeyConcurrency[apiKeyID] + } + return result, nil +} func (c *stubConcurrencyCacheForTest) IncrementWaitCount(_ context.Context, _ int64, _ int) (bool, error) { return c.waitAllowed, c.waitErr } @@ -201,6 +229,62 @@ func TestAcquireUserSlot_UnlimitedConcurrency(t *testing.T) { require.True(t, result.Acquired) } +func TestTrackAPIKeySlot_ReleaseDecrements(t *testing.T) { + cache := &stubConcurrencyCacheForTest{} + svc := NewConcurrencyService(cache) + + release := svc.TrackAPIKeySlot(context.Background(), 88) + require.NotNil(t, release) + require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs) + require.Len(t, cache.trackedAPIKeyRequestIDs, 1) + require.NotEmpty(t, cache.trackedAPIKeyRequestIDs[0]) + + release() + + require.Equal(t, []int64{88}, cache.releasedAPIKeyIDs) + require.Equal(t, cache.trackedAPIKeyRequestIDs, cache.releasedAPIKeyRequestIDs) +} + +func TestTrackAPIKeySlot_FailOpen(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyTrackErr: errors.New("redis down")} + svc := NewConcurrencyService(cache) + + release := svc.TrackAPIKeySlot(context.Background(), 88) + require.NotNil(t, release) + require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs) + + require.NotPanics(t, release) + require.Empty(t, cache.releasedAPIKeyIDs) +} + +func TestGetAPIKeyConcurrencyBatch_Fallbacks(t *testing.T) { + t.Run("nil cache returns zeroes", func(t *testing.T) { + svc := &ConcurrencyService{cache: nil} + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 0, 2: 0}, counts) + }) + + t.Run("redis error returns zeroes", func(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyConcurrencyErr: errors.New("redis down")} + svc := NewConcurrencyService(cache) + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 0, 2: 0}, counts) + }) + + t.Run("success returns counts", func(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyConcurrency: map[int64]int{1: 3, 2: 0}} + svc := NewConcurrencyService(cache) + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 3, 2: 0}, counts) + }) +} + func TestGenerateRequestID_UsesStablePrefixAndMonotonicCounter(t *testing.T) { id1 := generateRequestID() id2 := generateRequestID() diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index ab853c280a..bc1db19105 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -431,6 +431,20 @@ const ( // SettingKeyAllowUngroupedKeyScheduling 允许未分组 API Key 调度(默认 false:未分组 Key 返回 403) SettingKeyAllowUngroupedKeyScheduling = "allow_ungrouped_key_scheduling" + // SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled OpenAI 高级调度下是否启用粘性加权。 + SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled = "openai_advanced_scheduler_sticky_weighted_enabled" + // SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled OpenAI 高级调度下是否优先使用订阅账号池。 + SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled = "openai_advanced_scheduler_subscription_priority_enabled" + SettingKeyOpenAIAdvancedSchedulerLBTopK = "openai_advanced_scheduler_lb_top_k" + SettingKeyOpenAIAdvancedSchedulerWeightPriority = "openai_advanced_scheduler_weight_priority" + SettingKeyOpenAIAdvancedSchedulerWeightLoad = "openai_advanced_scheduler_weight_load" + SettingKeyOpenAIAdvancedSchedulerWeightQueue = "openai_advanced_scheduler_weight_queue" + SettingKeyOpenAIAdvancedSchedulerWeightErrorRate = "openai_advanced_scheduler_weight_error_rate" + SettingKeyOpenAIAdvancedSchedulerWeightTTFT = "openai_advanced_scheduler_weight_ttft" + SettingKeyOpenAIAdvancedSchedulerWeightReset = "openai_advanced_scheduler_weight_reset" + SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom = "openai_advanced_scheduler_weight_quota_headroom" + SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse = "openai_advanced_scheduler_weight_previous_response" + SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky = "openai_advanced_scheduler_weight_session_sticky" // SettingKeyBackendModeEnabled Backend 模式:禁用用户注册和自助服务,仅管理员可登录 SettingKeyBackendModeEnabled = "backend_mode_enabled" diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index f843ba3e45..35a60c8124 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -95,6 +95,9 @@ func (m *mockAccountRepoForPlatform) List(ctx context.Context, params pagination func (m *mockAccountRepoForPlatform) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForPlatform) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForPlatform) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/gateway_record_usage_test.go b/backend/internal/service/gateway_record_usage_test.go index c819eeca6e..2769251820 100644 --- a/backend/internal/service/gateway_record_usage_test.go +++ b/backend/internal/service/gateway_record_usage_test.go @@ -440,7 +440,9 @@ func TestGatewayServiceRecordUsage_GeneratesRequestIDWhenAllSourcesMissing(t *te require.Equal(t, billingRepo.lastCmd.RequestID, usageRepo.lastLog.RequestID) } -func TestGatewayServiceRecordUsage_DroppedUsageLogDoesNotSyncFallback(t *testing.T) { +func TestGatewayServiceRecordUsage_DroppedUsageLogFallsBackToSyncCreate(t *testing.T) { + // 计费成功后 best-effort 写入被丢弃(队列超时)时必须同步兜底, + // 否则出现“已扣费但无 usage_log”的对账缺口(issue #3656)。 usageRepo := &openAIRecordUsageBestEffortLogRepoStub{ bestEffortErr: MarkUsageLogCreateDropped(errors.New("usage log best-effort queue full")), } @@ -464,7 +466,9 @@ func TestGatewayServiceRecordUsage_DroppedUsageLogDoesNotSyncFallback(t *testing require.NoError(t, err) require.Equal(t, 1, usageRepo.bestEffortCalls) - require.Equal(t, 0, usageRepo.createCalls) + require.Equal(t, 1, usageRepo.createCalls) + // 兜底调用使用的 ctx 必须仍然存活,不能带着已死的 ctx 走过场。 + require.NoError(t, usageRepo.lastCtxErr) } func TestGatewayServiceRecordUsage_BillingErrorSkipsUsageLogWrite(t *testing.T) { diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 160a92a8e1..dcaf3a645c 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -5920,6 +5920,10 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( if c != nil && c.Request != nil { clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") } + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { body = sanitized } @@ -5956,6 +5960,9 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( setHeaderRaw(req.Header, "anthropic-version", "2023-06-01") } + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + return req, body, nil } @@ -6886,6 +6893,12 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex tokenType, mimicClaudeCode, modelID, clientHeaders, body, effectiveDropSet, ) + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值(由下方 ApplyHeaderOverrides 写入): + // body 能力净化必须以覆写值为准,否则 header/body 不对称会被上游 400。 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + finalBetaHeader, finalBetaShouldSet = beta, true + } + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { body = sanitized @@ -6959,6 +6972,10 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex } } + // 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op)。 + // 放在所有 header 逻辑之后,确保配置值对同名头拥有最终决定权。 + account.ApplyHeaderOverrides(req.Header) + // === DEBUG: 打印上游转发请求(headers + body 摘要),与 CLIENT_ORIGINAL 对比 === s.debugLogGatewaySnapshot("UPSTREAM_FORWARD", req.Header, body, map[string]string{ "url": req.URL.String(), @@ -9473,10 +9490,17 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage if writer, ok := repo.(usageLogBestEffortWriter); ok { if err := writer.CreateBestEffort(usageCtx, usageLog); err != nil { logger.LegacyPrintf(logKey, "Create usage log failed: %v", err) - if IsUsageLogCreateDropped(err) { - return + // 计费已在此前完成,日志必须落库:dropped(批处理队列超时)同样走同步兜底, + // 否则会出现“已扣费但无 usage_log”的对账缺口(issue #3656)。 + // 重复写入由 usage_logs 的 ON CONFLICT (request_id, api_key_id) DO NOTHING 防护。 + fallbackCtx := usageCtx + if usageCtx.Err() != nil { + // usageCtx 已耗尽(best-effort 入队阻塞到期限):换新的 detached 窗口,避免兜底必然失败。 + var fallbackCancel context.CancelFunc + fallbackCtx, fallbackCancel = detachedBillingContext(context.Background()) + defer fallbackCancel() } - if _, syncErr := repo.Create(usageCtx, usageLog); syncErr != nil { + if _, syncErr := repo.Create(fallbackCtx, usageLog); syncErr != nil { logger.LegacyPrintf(logKey, "Create usage log sync fallback failed: %v", syncErr) } } @@ -10403,6 +10427,10 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( if c != nil && c.Request != nil { clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") } + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { body = sanitized } @@ -10438,6 +10466,9 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( req.Header.Set("anthropic-version", "2023-06-01") } + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + return req, nil } @@ -10505,6 +10536,11 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con tokenType, mimicClaudeCode, modelID, clientHeaders, body, ctEffectiveDropSet, ) + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + finalBetaHeader, finalBetaShouldSet = beta, true + } + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { body = sanitized @@ -10571,6 +10607,9 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con } } + // 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + if c != nil && tokenType == "oauth" { c.Set(claudeMimicDebugInfoKey, buildClaudeMimicDebugLine(req, body, account, tokenType, mimicClaudeCode)) } diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index c021e88edf..7d5ed0ec9e 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -82,6 +82,9 @@ func (m *mockAccountRepoForGemini) List(ctx context.Context, params pagination.P func (m *mockAccountRepoForGemini) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForGemini) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForGemini) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 8942404eaa..f8b7ff3e83 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -483,6 +483,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody) case GrokMediaEndpointVideosGenerations: meta.ResponseID = extractGrokMediaVideoRequestID(responseBody) + // Video generation is one billable media unit; the legacy usage schema stores it in ImageCount. meta.ImageCount = 1 meta.ImageSize = requestInfo.SizeTier meta.ImageInputSize = requestInfo.Size diff --git a/backend/internal/service/group_capacity_service.go b/backend/internal/service/group_capacity_service.go index 459084dc59..b77b35773b 100644 --- a/backend/internal/service/group_capacity_service.go +++ b/backend/internal/service/group_capacity_service.go @@ -16,6 +16,26 @@ type GroupCapacitySummary struct { RPMMax int `json:"rpm_max"` } +// GroupAccountCapacityRow is the lightweight account projection needed for +// capacity summary aggregation. +type GroupAccountCapacityRow struct { + GroupID int64 + AccountID int64 + Concurrency int + Extra map[string]any + SessionWindowStart *time.Time + SessionWindowEnd *time.Time + SessionWindowStatus string +} + +type groupCapacityActiveGroupIDLister interface { + ListActiveIDs(ctx context.Context) ([]int64, error) +} + +type groupCapacityAccountLister interface { + ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error) +} + // GroupCapacityService aggregates per-group capacity from runtime data. type GroupCapacityService struct { accountRepo AccountRepository @@ -44,24 +64,176 @@ func NewGroupCapacityService( // GetAllGroupCapacity returns capacity summary for all active groups. func (s *GroupCapacityService) GetAllGroupCapacity(ctx context.Context) ([]GroupCapacitySummary, error) { - groups, err := s.groupRepo.ListActive(ctx) + groupIDs, err := s.listActiveGroupIDs(ctx) if err != nil { return nil, err } - results := make([]GroupCapacitySummary, 0, len(groups)) + if lister, ok := s.accountRepo.(groupCapacityAccountLister); ok { + return s.getGroupCapacitiesBatch(ctx, groupIDs, lister) + } + + return s.getGroupCapacitiesSequential(ctx, groupIDs), nil +} + +func (s *GroupCapacityService) listActiveGroupIDs(ctx context.Context) ([]int64, error) { + if lister, ok := s.groupRepo.(groupCapacityActiveGroupIDLister); ok { + return lister.ListActiveIDs(ctx) + } + + groups, err := s.groupRepo.ListActive(ctx) + if err != nil { + return nil, err + } + groupIDs := make([]int64, 0, len(groups)) for i := range groups { - cap, err := s.getGroupCapacity(ctx, groups[i].ID) + groupIDs = append(groupIDs, groups[i].ID) + } + return groupIDs, nil +} + +func (s *GroupCapacityService) getGroupCapacitiesSequential(ctx context.Context, groupIDs []int64) []GroupCapacitySummary { + results := make([]GroupCapacitySummary, 0, len(groupIDs)) + for _, groupID := range groupIDs { + cap, err := s.getGroupCapacity(ctx, groupID) if err != nil { // Skip groups with errors, return partial results continue } - cap.GroupID = groups[i].ID + cap.GroupID = groupID results = append(results, cap) } + return results +} + +type groupCapacityAccountRef struct { + groupID int64 + accountID int64 +} + +func (s *GroupCapacityService) getGroupCapacitiesBatch(ctx context.Context, groupIDs []int64, lister groupCapacityAccountLister) ([]GroupCapacitySummary, error) { + results := make([]GroupCapacitySummary, len(groupIDs)) + groupIndex := make(map[int64]int, len(groupIDs)) + for i, groupID := range groupIDs { + results[i].GroupID = groupID + groupIndex[groupID] = i + } + if len(groupIDs) == 0 { + return results, nil + } + + rows, err := lister.ListSchedulableCapacityByGroupIDs(ctx, groupIDs) + if err != nil { + return nil, err + } + if len(rows) == 0 { + return results, nil + } + + refs := make([]groupCapacityAccountRef, 0, len(rows)) + seenGroupAccount := make(map[groupCapacityAccountRef]struct{}, len(rows)) + accountIDSet := make(map[int64]struct{}, len(rows)) + accountIDs := make([]int64, 0, len(rows)) + sessionTimeouts := make(map[int64]time.Duration) + + for _, row := range rows { + idx, ok := groupIndex[row.GroupID] + if !ok || row.AccountID <= 0 { + continue + } + + ref := groupCapacityAccountRef{groupID: row.GroupID, accountID: row.AccountID} + if _, ok := seenGroupAccount[ref]; ok { + continue + } + seenGroupAccount[ref] = struct{}{} + refs = append(refs, ref) + + if _, ok := accountIDSet[row.AccountID]; !ok { + accountIDSet[row.AccountID] = struct{}{} + accountIDs = append(accountIDs, row.AccountID) + } + + acc := Account{ + ID: row.AccountID, + Concurrency: row.Concurrency, + Extra: row.Extra, + SessionWindowStart: row.SessionWindowStart, + SessionWindowEnd: row.SessionWindowEnd, + SessionWindowStatus: row.SessionWindowStatus, + } + + results[idx].ConcurrencyMax += acc.Concurrency + + if maxSessions := acc.GetMaxSessions(); maxSessions > 0 { + results[idx].SessionsMax += maxSessions + timeout := time.Duration(acc.GetSessionIdleTimeoutMinutes()) * time.Minute + if timeout <= 0 { + timeout = 5 * time.Minute + } + sessionTimeouts[acc.ID] = timeout + } + + if rpm := acc.GetBaseRPM(); rpm > 0 { + results[idx].RPMMax += rpm + } + } + + if len(accountIDs) == 0 { + return results, nil + } + + concurrencyMap := map[int64]int{} + if s.concurrencyService != nil { + concurrencyMap, _ = s.concurrencyService.GetAccountConcurrencyBatch(ctx, accountIDs) + } + + sessionAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool { + return summary.SessionsMax > 0 + }) + var sessionsMap map[int64]int + if len(sessionAccountIDs) > 0 && s.sessionLimitCache != nil { + sessionsMap, _ = s.sessionLimitCache.GetActiveSessionCountBatch(ctx, sessionAccountIDs, sessionTimeouts) + } + + rpmAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool { + return summary.RPMMax > 0 + }) + var rpmMap map[int64]int + if len(rpmAccountIDs) > 0 && s.rpmCache != nil { + rpmMap, _ = s.rpmCache.GetRPMBatch(ctx, rpmAccountIDs) + } + + for _, ref := range refs { + idx := groupIndex[ref.groupID] + results[idx].ConcurrencyUsed += concurrencyMap[ref.accountID] + if sessionsMap != nil && results[idx].SessionsMax > 0 { + results[idx].SessionsUsed += sessionsMap[ref.accountID] + } + if rpmMap != nil && results[idx].RPMMax > 0 { + results[idx].RPMUsed += rpmMap[ref.accountID] + } + } return results, nil } +func accountIDsForGroupsWithLimit(refs []groupCapacityAccountRef, groupIndex map[int64]int, summaries []GroupCapacitySummary, include func(GroupCapacitySummary) bool) []int64 { + seen := make(map[int64]struct{}) + accountIDs := make([]int64, 0) + for _, ref := range refs { + idx, ok := groupIndex[ref.groupID] + if !ok || !include(summaries[idx]) { + continue + } + if _, ok := seen[ref.accountID]; ok { + continue + } + seen[ref.accountID] = struct{}{} + accountIDs = append(accountIDs, ref.accountID) + } + return accountIDs +} + func (s *GroupCapacityService) getGroupCapacity(ctx context.Context, groupID int64) (GroupCapacitySummary, error) { accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, groupID) if err != nil { diff --git a/backend/internal/service/group_capacity_service_test.go b/backend/internal/service/group_capacity_service_test.go new file mode 100644 index 0000000000..73927307d2 --- /dev/null +++ b/backend/internal/service/group_capacity_service_test.go @@ -0,0 +1,179 @@ +package service + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type groupCapacityAccountRepoStub struct { + AccountRepository + rows []GroupAccountCapacityRow + requested []int64 +} + +func (s *groupCapacityAccountRepoStub) ListSchedulableCapacityByGroupIDs(_ context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error) { + s.requested = append([]int64(nil), groupIDs...) + return append([]GroupAccountCapacityRow(nil), s.rows...), nil +} + +type groupCapacityGroupRepoStub struct { + GroupRepository + groupIDs []int64 + listCalls int +} + +func (s *groupCapacityGroupRepoStub) ListActiveIDs(context.Context) ([]int64, error) { + s.listCalls++ + return append([]int64(nil), s.groupIDs...), nil +} + +type groupCapacityConcurrencyCacheStub struct { + ConcurrencyCache + counts map[int64]int + requested []int64 +} + +func (s *groupCapacityConcurrencyCacheStub) GetAccountConcurrencyBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) { + s.requested = append([]int64(nil), accountIDs...) + out := make(map[int64]int, len(accountIDs)) + for _, id := range accountIDs { + out[id] = s.counts[id] + } + return out, nil +} + +type groupCapacitySessionCacheStub struct { + SessionLimitCache + counts map[int64]int + requested []int64 + idleTimeouts map[int64]time.Duration +} + +func (s *groupCapacitySessionCacheStub) GetActiveSessionCountBatch(_ context.Context, accountIDs []int64, idleTimeouts map[int64]time.Duration) (map[int64]int, error) { + s.requested = append([]int64(nil), accountIDs...) + s.idleTimeouts = make(map[int64]time.Duration, len(idleTimeouts)) + for id, timeout := range idleTimeouts { + s.idleTimeouts[id] = timeout + } + out := make(map[int64]int, len(accountIDs)) + for _, id := range accountIDs { + out[id] = s.counts[id] + } + return out, nil +} + +type groupCapacityRPMCacheStub struct { + RPMCache + counts map[int64]int + requested []int64 +} + +func (s *groupCapacityRPMCacheStub) GetRPMBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) { + s.requested = append([]int64(nil), accountIDs...) + out := make(map[int64]int, len(accountIDs)) + for _, id := range accountIDs { + out[id] = s.counts[id] + } + return out, nil +} + +func TestGetAllGroupCapacityBatchAggregatesRuntimeAndLimits(t *testing.T) { + accountRepo := &groupCapacityAccountRepoStub{ + rows: []GroupAccountCapacityRow{ + { + GroupID: 10, + AccountID: 1, + Concurrency: 2, + Extra: map[string]any{ + "max_sessions": 3, + "session_idle_timeout_minutes": 7, + "base_rpm": 11, + }, + }, + { + GroupID: 20, + AccountID: 1, + Concurrency: 2, + Extra: map[string]any{ + "max_sessions": 3, + "session_idle_timeout_minutes": 7, + "base_rpm": 11, + }, + }, + { + GroupID: 20, + AccountID: 2, + Concurrency: 4, + Extra: map[string]any{ + "max_sessions": 1, + "session_idle_timeout_minutes": 9, + "base_rpm": 13, + }, + }, + }, + } + groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}} + concurrencyCache := &groupCapacityConcurrencyCacheStub{counts: map[int64]int{1: 1, 2: 2}} + sessionCache := &groupCapacitySessionCacheStub{counts: map[int64]int{1: 2, 2: 1}} + rpmCache := &groupCapacityRPMCacheStub{counts: map[int64]int{1: 5, 2: 7}} + svc := NewGroupCapacityService( + accountRepo, + groupRepo, + NewConcurrencyService(concurrencyCache), + sessionCache, + rpmCache, + ) + + results, err := svc.GetAllGroupCapacity(context.Background()) + require.NoError(t, err) + + require.Equal(t, 1, groupRepo.listCalls) + require.Equal(t, []int64{10, 20}, accountRepo.requested) + require.Equal(t, []int64{1, 2}, concurrencyCache.requested) + require.ElementsMatch(t, []int64{1, 2}, sessionCache.requested) + require.ElementsMatch(t, []int64{1, 2}, rpmCache.requested) + require.Equal(t, 7*time.Minute, sessionCache.idleTimeouts[1]) + require.Equal(t, 9*time.Minute, sessionCache.idleTimeouts[2]) + + require.Equal(t, []GroupCapacitySummary{ + { + GroupID: 10, + ConcurrencyUsed: 1, + ConcurrencyMax: 2, + SessionsUsed: 2, + SessionsMax: 3, + RPMUsed: 5, + RPMMax: 11, + }, + { + GroupID: 20, + ConcurrencyUsed: 3, + ConcurrencyMax: 6, + SessionsUsed: 3, + SessionsMax: 4, + RPMUsed: 12, + RPMMax: 24, + }, + }, results) +} + +func TestGetAllGroupCapacityBatchKeepsEmptyGroupRows(t *testing.T) { + accountRepo := &groupCapacityAccountRepoStub{ + rows: []GroupAccountCapacityRow{ + {GroupID: 20, AccountID: 2, Concurrency: 4}, + }, + } + groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}} + svc := NewGroupCapacityService(accountRepo, groupRepo, nil, nil, nil) + + results, err := svc.GetAllGroupCapacity(context.Background()) + require.NoError(t, err) + + require.Equal(t, []GroupCapacitySummary{ + {GroupID: 10}, + {GroupID: 20, ConcurrencyMax: 4}, + }, results) +} diff --git a/backend/internal/service/model_rate_limit.go b/backend/internal/service/model_rate_limit.go index 1568a00f78..195f5a1dfd 100644 --- a/backend/internal/service/model_rate_limit.go +++ b/backend/internal/service/model_rate_limit.go @@ -12,6 +12,9 @@ const ( modelRateLimitsKey = "model_rate_limits" antigravityGeminiModelRateLimitKey = "antigravity:gemini" openAIImageGenerationRateLimitKey = "openai:image_generation" + // anthropicFableRateLimitKey 是 Anthropic 7d_oi(Fable 专属 7d 窗口)限流的 + // 家族级 scope:命中后所有 Fable 变体(含 [1m] 等后缀)都不再调度到该账号。 + anthropicFableRateLimitKey = "claude-fable-5" ) // isRateLimitActiveForKey 检查指定 key 的限流是否生效 @@ -82,10 +85,19 @@ func (a *Account) modelRateLimitKeysForRequest(ctx context.Context, requestedMod if openAIImageGenerationRateLimitApplies(ctx, requestedModel, modelKey) && modelKey != openAIImageGenerationRateLimitKey { keys = append(keys, openAIImageGenerationRateLimitKey) } + case PlatformAnthropic: + if isAnthropicFableModel(modelKey) && modelKey != anthropicFableRateLimitKey { + keys = append(keys, anthropicFableRateLimitKey) + } } return keys } +// isAnthropicFableModel 判断是否为 Fable 模型家族(claude-fable-5、claude-fable-5[1m] 等变体) +func isAnthropicFableModel(model string) bool { + return strings.Contains(strings.ToLower(model), "fable") +} + func openAIImageGenerationRateLimitApplies(ctx context.Context, requestedModel, modelKey string) bool { if isOpenAIImageGenerationModel(requestedModel) || isOpenAIImageGenerationModel(modelKey) { return true diff --git a/backend/internal/service/model_rate_limit_test.go b/backend/internal/service/model_rate_limit_test.go index c62e4c4b9b..430d91baf6 100644 --- a/backend/internal/service/model_rate_limit_test.go +++ b/backend/internal/service/model_rate_limit_test.go @@ -499,3 +499,47 @@ func TestGetRateLimitRemainingTime(t *testing.T) { }) } } + +func TestIsModelRateLimited_AnthropicFableFamilyKey(t *testing.T) { + now := time.Now() + future := now.Add(48 * time.Hour).Format(time.RFC3339) + + account := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + modelRateLimitsKey: map[string]any{ + anthropicFableRateLimitKey: map[string]any{ + "rate_limit_reset_at": future, + }, + }, + }, + } + + tests := []struct { + requestedModel string + expected bool + }{ + {"claude-fable-5", true}, + {"claude-fable-5[1m]", true}, // 家族 key 覆盖变体 + {"Claude-Fable-5-20260601", true}, // 大小写不敏感 + {"claude-sonnet-4-6", false}, // 其他模型不受影响 + {"claude-opus-4-8", false}, + } + + for _, tc := range tests { + t.Run(tc.requestedModel, func(t *testing.T) { + got := account.isModelRateLimitedWithContext(context.Background(), tc.requestedModel) + require.Equal(t, tc.expected, got) + remaining := account.GetModelRateLimitRemainingTimeWithContext(context.Background(), tc.requestedModel) + require.Equal(t, tc.expected, remaining > 0) + }) + } +} + +func TestIsAnthropicFableModel(t *testing.T) { + require.True(t, isAnthropicFableModel("claude-fable-5")) + require.True(t, isAnthropicFableModel("claude-fable-5[1m]")) + require.True(t, isAnthropicFableModel("Claude-Fable-5")) + require.False(t, isAnthropicFableModel("claude-sonnet-4-6")) + require.False(t, isAnthropicFableModel("")) +} diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index ab298d4521..a0a4fbff3e 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -36,8 +36,20 @@ const ( ) type cachedOpenAIAdvancedSchedulerSetting struct { - enabled bool - expiresAt int64 + enabled bool + stickyWeightedEnabled bool + subscriptionPriorityEnabled bool + lbTopKOverride int + weightOverrides map[string]float64 + expiresAt int64 +} + +type openAIAdvancedSchedulerRuntimeSettings struct { + enabled bool + stickyWeightedEnabled bool + subscriptionPriorityEnabled bool + lbTopKOverride int + weightOverrides map[string]float64 } var openAIAdvancedSchedulerSettingCache atomic.Value // *cachedOpenAIAdvancedSchedulerSetting @@ -48,8 +60,12 @@ type OpenAIAccountScheduleRequest struct { Platform string SessionHash string StickyAccountID int64 + StickyPreviousAccountID int64 + StickyWeighted bool + SubscriptionPriority bool PreserveStickyBinding bool PreviousResponseID string + PreviousResponseCanMove bool RequestedModel string RequiredTransport OpenAIUpstreamTransport RequiredCapability OpenAIEndpointCapability @@ -111,6 +127,17 @@ type openAIAccountLoadPlan struct { loadSkew float64 } +type openAIAccountLoadSelectionAttempt struct { + result *AccountSelectionResult + selectionOrder []openAIAccountCandidateScore + candidateCount int + topK int + loadSkew float64 + compactBlocked bool + noCompactCandidates bool + err error +} + func (m *openAIAccountSchedulerMetrics) recordSelect(decision OpenAIAccountScheduleDecision) { if m == nil { return @@ -277,7 +304,8 @@ func (s *defaultOpenAIAccountScheduler) Select( }() previousResponseID := strings.TrimSpace(req.PreviousResponseID) - if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI { + if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI && + (!req.StickyWeighted || !req.PreviousResponseCanMove) { selection, err := s.service.selectAccountByPreviousResponseIDForCapability( ctx, req.GroupID, @@ -310,19 +338,21 @@ func (s *defaultOpenAIAccountScheduler) Select( } } - selection, escapedSticky, err := s.selectBySessionHash(ctx, req) - if err != nil { - return nil, decision, err - } - if selection != nil && selection.Account != nil { - decision.Layer = openAIAccountScheduleLayerSessionSticky - decision.StickySessionHit = true - decision.SelectedAccountID = selection.Account.ID - decision.SelectedAccountType = selection.Account.Type - return selection, decision, nil - } - if escapedSticky { - req.PreserveStickyBinding = true + if !req.StickyWeighted { + selection, escapedSticky, err := s.selectBySessionHash(ctx, req) + if err != nil { + return nil, decision, err + } + if selection != nil && selection.Account != nil { + decision.Layer = openAIAccountScheduleLayerSessionSticky + decision.StickySessionHit = true + decision.SelectedAccountID = selection.Account.ID + decision.SelectedAccountType = selection.Account.Type + return selection, decision, nil + } + if escapedSticky { + req.PreserveStickyBinding = true + } } selection, candidateCount, topK, loadSkew, err := s.selectByLoadBalance(ctx, req) @@ -336,6 +366,14 @@ func (s *defaultOpenAIAccountScheduler) Select( if selection != nil && selection.Account != nil { decision.SelectedAccountID = selection.Account.ID decision.SelectedAccountType = selection.Account.Type + if req.StickyWeighted { + if req.StickyPreviousAccountID > 0 && selection.Account.ID == req.StickyPreviousAccountID { + decision.StickyPreviousHit = true + } + if req.StickyAccountID > 0 && selection.Account.ID == req.StickyAccountID { + decision.StickySessionHit = true + } + } } return selection, decision, nil } @@ -453,6 +491,13 @@ func openAIStickyAccountMatchesGroup(account *Account, groupID *int64) bool { return false } +func openAIAccountSchedulingPriority(account *Account) int { + if account == nil { + return 0 + } + return account.Priority +} + func (s *defaultOpenAIAccountScheduler) shouldEscapeStickyAccount(accountID int64, cfg openAIStickyEscapeConfig) (reason string, errorRate float64, ttft float64, shouldEscape bool) { if !cfg.enabled || s == nil || s.stats == nil || accountID <= 0 { return "", 0, 0, false @@ -471,6 +516,7 @@ type openAIAccountCandidateScore struct { account *Account loadInfo *AccountLoadInfo score float64 + priority int errorRate float64 ttft float64 hasTTFT bool @@ -669,6 +715,7 @@ func buildOpenAIWeightedSelectionOrder( } func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan( + ctx context.Context, req OpenAIAccountScheduleRequest, filtered []*Account, loadMap map[int64]*AccountLoadInfo, @@ -716,18 +763,20 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan( return plan } - minPriority, maxPriority := candidates[0].account.Priority, candidates[0].account.Priority + minPriority, maxPriority := openAIAccountSchedulingPriority(candidates[0].account), openAIAccountSchedulingPriority(candidates[0].account) maxWaiting := 1 loadRateSum := 0.0 loadRateSumSquares := 0.0 minTTFT, maxTTFT := 0.0, 0.0 hasTTFTSample := false - for _, candidate := range candidates { - if candidate.account.Priority < minPriority { - minPriority = candidate.account.Priority + for i := range candidates { + candidate := &candidates[i] + candidate.priority = openAIAccountSchedulingPriority(candidate.account) + if candidate.priority < minPriority { + minPriority = candidate.priority } - if candidate.account.Priority > maxPriority { - maxPriority = candidate.account.Priority + if candidate.priority > maxPriority { + maxPriority = candidate.priority } if candidate.loadInfo.WaitingCount > maxWaiting { maxWaiting = candidate.loadInfo.WaitingCount @@ -751,7 +800,7 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan( } plan.loadSkew = calcLoadSkewByMoments(loadRateSum, loadRateSumSquares, len(candidates)) - weights := s.service.openAIWSSchedulerWeights() + weights := s.service.openAIWSSchedulerWeightsForRequest(ctx) // Reset 因子(use-it-or-lose-it):在拥有「未来会话窗口结束时间」的账号中, // 剩余时间越短 → 因子越接近 1(越早重置越优先用尽)。无活跃窗口的账号因子为 0。 @@ -785,7 +834,7 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan( item := &candidates[i] priorityFactor := 1.0 if maxPriority > minPriority { - priorityFactor = 1 - float64(item.account.Priority-minPriority)/float64(maxPriority-minPriority) + priorityFactor = 1 - float64(item.priority-minPriority)/float64(maxPriority-minPriority) } loadFactor := 1 - clamp01(float64(item.loadInfo.LoadRate)/100.0) queueFactor := 1 - clamp01(float64(item.loadInfo.WaitingCount)/float64(maxWaiting)) @@ -817,10 +866,18 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan( weights.TTFT*ttftFactor + weights.Reset*resetFactor + weights.QuotaHeadroom*quotaHeadroomFactor + if req.StickyWeighted { + if req.PreviousResponseCanMove && req.StickyPreviousAccountID > 0 && item.account.ID == req.StickyPreviousAccountID { + item.score += weights.Previous + } + if req.StickyAccountID > 0 && item.account.ID == req.StickyAccountID { + item.score += weights.SessionSticky + } + } } plan.candidates = candidates - plan.topK = s.service.openAIWSLBTopK() + plan.topK = s.service.openAIWSLBTopKForRequest(ctx) if plan.topK > len(candidates) { plan.topK = len(candidates) } @@ -845,6 +902,20 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAISelectionOrder( groupTopK = len(pool) } ranked := selectTopKOpenAICandidates(pool, groupTopK) + if req.StickyWeighted { + for _, stickyID := range []int64{req.StickyPreviousAccountID, req.StickyAccountID} { + if stickyID <= 0 { + continue + } + for i, candidate := range ranked { + if candidate.account != nil && candidate.account.ID == stickyID { + ordered := append([]openAIAccountCandidateScore{candidate}, ranked[:i]...) + ordered = append(ordered, ranked[i+1:]...) + return ordered + } + } + } + } return buildOpenAIWeightedSelectionOrder(ranked, req) } @@ -939,6 +1010,74 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder( return nil, compactBlocked, nil } +func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( + ctx context.Context, + req OpenAIAccountScheduleRequest, +) (*AccountSelectionResult, error) { + if !req.StickyWeighted { + return nil, nil + } + for _, accountID := range []int64{req.StickyPreviousAccountID, req.StickyAccountID} { + if accountID <= 0 { + continue + } + if req.ExcludedIDs != nil { + if _, excluded := req.ExcludedIDs[accountID]; excluded { + continue + } + } + account, err := s.service.getSchedulableAccount(ctx, accountID) + if err != nil || account == nil { + continue + } + if !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) { + continue + } + account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.Platform, req.RequestedModel, req.RequireCompact, req.RequiredCapability) + if account == nil || !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) { + continue + } + // 粘性绑定只证明绑定时账号在分组内;账号被移出分组后绑定仍会在 TTL 内存活, + // 必须与 selectBySessionHash 一样重验分组归属,否则会把分组流量泄漏到组外账号。 + if !openAIStickyAccountMatchesGroup(account, req.GroupID) { + if accountID == req.StickyAccountID && strings.TrimSpace(req.SessionHash) != "" { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, req.SessionHash) + } + continue + } + if req.RequireCompact && openAICompactSupportTier(account) == 0 { + continue + } + result, acquireErr := s.service.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) + if acquireErr != nil { + return nil, acquireErr + } + if result != nil && result.Acquired { + if req.SessionHash != "" && !req.PreserveStickyBinding { + _ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, account.ID) + } + return &AccountSelectionResult{ + Account: account, + Acquired: true, + ReleaseFunc: result.ReleaseFunc, + }, nil + } + if s.service.concurrencyService != nil { + cfg := s.service.schedulingConfig() + return &AccountSelectionResult{ + Account: account, + WaitPlan: &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }, + }, nil + } + } + return nil, nil +} + func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( ctx context.Context, req OpenAIAccountScheduleRequest, @@ -1002,52 +1141,175 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( } } - plan := s.buildOpenAIAccountLoadPlan(req, filtered, loadMap) - candidateCount := plan.candidateCount - topK := plan.topK - loadSkew := plan.loadSkew - selectionOrder := plan.selectionOrder - if req.RequireCompact && len(plan.candidates) == 0 && len(plan.staleSnapshotCompactRetry) == 0 { - return nil, 0, 0, 0, ErrNoAvailableCompactAccounts - } - if req.RequireCompact && len(selectionOrder) == 0 && s.service.schedulerSnapshot == nil { - return nil, candidateCount, topK, loadSkew, ErrNoAvailableCompactAccounts - } - if len(selectionOrder) == 0 { - return nil, candidateCount, topK, loadSkew, noAvailableOpenAISelectionError(req.RequestedModel, req.RequireCompact && len(plan.allCandidates) > 0) + if req.SubscriptionPriority { + subscriptionAccounts, regularAccounts := partitionOpenAIChatGPTSubscriptionAccounts(filtered) + if len(subscriptionAccounts) > 0 { + attempt := s.trySelectByLoadBalancePool(ctx, req, subscriptionAccounts, loadMap) + if attempt.err != nil && (!attempt.noCompactCandidates || len(regularAccounts) <= 0) { + return nil, attempt.candidateCount, attempt.topK, attempt.loadSkew, attempt.err + } + if attempt.result != nil { + return attempt.result, attempt.candidateCount, attempt.topK, attempt.loadSkew, nil + } + if len(regularAccounts) > 0 { + regularAttempt := s.trySelectByLoadBalancePool(ctx, req, regularAccounts, loadMap) + if regularAttempt.err != nil && !regularAttempt.noCompactCandidates { + return nil, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, regularAttempt.err + } + if regularAttempt.result != nil { + return regularAttempt.result, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, nil + } + var result *AccountSelectionResult + candidateCount, topK, loadSkew := regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew + fallbackErr := regularAttempt.err + if regularAttempt.err == nil { + result, candidateCount, topK, loadSkew, fallbackErr = s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt) + if fallbackErr == nil && result != nil { + return result, candidateCount, topK, loadSkew, nil + } + } + // 常规池既无法获取也无法排队(含仅剩不支持 compact 的候选)时, + // 回退到订阅池的等待计划:busy-but-waitable 的订阅账号不应因常规池存在 + // 而被丢弃,否则开启订阅优先反而让本可排队成功的请求硬失败。 + subResult, subCandidateCount, subTopK, subLoadSkew, subErr := s.finishLoadBalanceSelectionFallback(ctx, req, attempt) + if subErr == nil && subResult != nil { + return subResult, subCandidateCount, subTopK, subLoadSkew, nil + } + return result, candidateCount, topK, loadSkew, fallbackErr + } + return s.finishLoadBalanceSelectionFallback(ctx, req, attempt) + } } - result, compactBlocked, acquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, selectionOrder) + attempt := s.trySelectByLoadBalancePool(ctx, req, filtered, loadMap) + if attempt.err != nil { + return nil, attempt.candidateCount, attempt.topK, attempt.loadSkew, attempt.err + } + if attempt.result != nil { + return attempt.result, attempt.candidateCount, attempt.topK, attempt.loadSkew, nil + } + return s.finishLoadBalanceSelectionFallback(ctx, req, attempt) +} + +func partitionOpenAIChatGPTSubscriptionAccounts(accounts []*Account) ([]*Account, []*Account) { + subscriptionAccounts := make([]*Account, 0, len(accounts)) + regularAccounts := make([]*Account, 0, len(accounts)) + for _, account := range accounts { + if account != nil && account.IsOpenAIChatGPTSubscription() { + subscriptionAccounts = append(subscriptionAccounts, account) + continue + } + regularAccounts = append(regularAccounts, account) + } + return subscriptionAccounts, regularAccounts +} + +func (s *defaultOpenAIAccountScheduler) trySelectByLoadBalancePool( + ctx context.Context, + req OpenAIAccountScheduleRequest, + filtered []*Account, + loadMap map[int64]*AccountLoadInfo, +) openAIAccountLoadSelectionAttempt { + plan := s.buildOpenAIAccountLoadPlan(ctx, req, filtered, loadMap) + attempt := openAIAccountLoadSelectionAttempt{ + selectionOrder: plan.selectionOrder, + candidateCount: plan.candidateCount, + topK: plan.topK, + loadSkew: plan.loadSkew, + } + if req.RequireCompact && len(plan.candidates) == 0 && len(plan.staleSnapshotCompactRetry) == 0 { + attempt.noCompactCandidates = true + attempt.err = ErrNoAvailableCompactAccounts + return attempt + } + if req.RequireCompact && len(attempt.selectionOrder) == 0 && s.service.schedulerSnapshot == nil { + attempt.noCompactCandidates = true + attempt.err = ErrNoAvailableCompactAccounts + return attempt + } + if len(attempt.selectionOrder) == 0 { + attempt.compactBlocked = req.RequireCompact && len(plan.allCandidates) > 0 + return attempt + } + + result, compactBlocked, acquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, attempt.selectionOrder) + attempt.compactBlocked = compactBlocked if acquireErr != nil { - return nil, candidateCount, topK, loadSkew, acquireErr + attempt.err = acquireErr + return attempt } if result != nil { - return result, candidateCount, topK, loadSkew, nil + attempt.result = result + return attempt } if s.service.concurrencyService != nil { + loadReq := buildOpenAIAccountLoadRequest(filtered) if freshLoadMap, loadErr := s.service.concurrencyService.GetAccountsLoadBatchFresh(ctx, loadReq); loadErr == nil { - freshPlan := s.buildOpenAIAccountLoadPlan(req, filtered, freshLoadMap) + freshPlan := s.buildOpenAIAccountLoadPlan(ctx, req, filtered, freshLoadMap) if len(freshPlan.selectionOrder) > 0 { freshResult, freshCompactBlocked, freshAcquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, freshPlan.selectionOrder) if freshAcquireErr != nil { - return nil, candidateCount, topK, loadSkew, freshAcquireErr + attempt.err = freshAcquireErr + return attempt } if freshResult != nil { - return freshResult, freshPlan.candidateCount, freshPlan.topK, freshPlan.loadSkew, nil + attempt.result = freshResult + attempt.selectionOrder = freshPlan.selectionOrder + attempt.candidateCount = freshPlan.candidateCount + attempt.topK = freshPlan.topK + attempt.loadSkew = freshPlan.loadSkew + return attempt } - compactBlocked = compactBlocked || freshCompactBlocked - selectionOrder = freshPlan.selectionOrder - candidateCount = freshPlan.candidateCount - topK = freshPlan.topK - loadSkew = freshPlan.loadSkew + attempt.compactBlocked = attempt.compactBlocked || freshCompactBlocked + attempt.selectionOrder = freshPlan.selectionOrder + attempt.candidateCount = freshPlan.candidateCount + attempt.topK = freshPlan.topK + attempt.loadSkew = freshPlan.loadSkew } } } + return attempt +} + +func buildOpenAIAccountLoadRequest(accounts []*Account) []AccountWithConcurrency { + loadReq := make([]AccountWithConcurrency, 0, len(accounts)) + for _, account := range accounts { + if account == nil { + continue + } + loadReq = append(loadReq, AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + return loadReq +} + +func (s *defaultOpenAIAccountScheduler) finishLoadBalanceSelectionFallback( + ctx context.Context, + req OpenAIAccountScheduleRequest, + attempt openAIAccountLoadSelectionAttempt, +) (*AccountSelectionResult, int, int, float64, error) { + candidateCount := attempt.candidateCount + topK := attempt.topK + loadSkew := attempt.loadSkew + + if len(attempt.selectionOrder) == 0 { + return nil, candidateCount, topK, loadSkew, noAvailableOpenAISelectionError(req.RequestedModel, attempt.compactBlocked) + } + + if stickyFallback, stickyErr := s.tryFallbackToWeightedSticky(ctx, req); stickyErr != nil { + return nil, candidateCount, topK, loadSkew, stickyErr + } else if stickyFallback != nil { + return stickyFallback, candidateCount, topK, loadSkew, nil + } + cfg := s.service.schedulingConfig() + compactBlocked := attempt.compactBlocked // WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。 - for _, candidate := range selectionOrder { + for _, candidate := range attempt.selectionOrder { fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.Platform, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue @@ -1184,40 +1446,169 @@ func (s *OpenAIGatewayService) openAIAdvancedSchedulerSettingRepo() SettingRepos return s.rateLimitService.settingService.settingRepo } -func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerEnabled(ctx context.Context) bool { +func (s *OpenAIGatewayService) openAIAdvancedSchedulerRuntimeSettings(ctx context.Context) openAIAdvancedSchedulerRuntimeSettings { if cached, ok := openAIAdvancedSchedulerSettingCache.Load().(*cachedOpenAIAdvancedSchedulerSetting); ok && cached != nil { if time.Now().UnixNano() < cached.expiresAt { - return cached.enabled + return openAIAdvancedSchedulerRuntimeSettings{ + enabled: cached.enabled, + stickyWeightedEnabled: cached.stickyWeightedEnabled, + subscriptionPriorityEnabled: cached.subscriptionPriorityEnabled, + lbTopKOverride: cached.lbTopKOverride, + weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(cached.weightOverrides), + } } } result, _, _ := openAIAdvancedSchedulerSettingSF.Do(openAIAdvancedSchedulerSettingKey, func() (any, error) { if cached, ok := openAIAdvancedSchedulerSettingCache.Load().(*cachedOpenAIAdvancedSchedulerSetting); ok && cached != nil { if time.Now().UnixNano() < cached.expiresAt { - return cached.enabled, nil + return openAIAdvancedSchedulerRuntimeSettings{ + enabled: cached.enabled, + stickyWeightedEnabled: cached.stickyWeightedEnabled, + subscriptionPriorityEnabled: cached.subscriptionPriorityEnabled, + lbTopKOverride: cached.lbTopKOverride, + weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(cached.weightOverrides), + }, nil } } enabled := false + stickyWeightedEnabled := false + subscriptionPriorityEnabled := false + lbTopKOverride := 0 + weightOverrides := map[string]float64{} if repo := s.openAIAdvancedSchedulerSettingRepo(); repo != nil { dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIAdvancedSchedulerSettingDBTimeout) defer cancel() - value, err := repo.GetValue(dbCtx, openAIAdvancedSchedulerSettingKey) - if err == nil { - enabled = strings.EqualFold(strings.TrimSpace(value), "true") + if values, err := repo.GetMultiple(dbCtx, openAIAdvancedSchedulerRuntimeSettingKeys()); err == nil { + enabled = strings.EqualFold(strings.TrimSpace(values[openAIAdvancedSchedulerSettingKey]), "true") + stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(values[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true") + subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(values[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true") + lbTopKOverride = parsePositiveIntOverride(values[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(values) + } else { + // 批量读取失败时逐键降级,覆盖全部键(含 TopK/权重),避免只加载布尔开关 + // 而静默丢弃管理员配置的覆盖值;降级状态会被缓存一个 TTL,必须留痕。 + slog.Warn("openai_advanced_scheduler_settings_batch_load_failed", "error", err) + fallbackValues := make(map[string]string) + for _, key := range openAIAdvancedSchedulerRuntimeSettingKeys() { + if value, valueErr := repo.GetValue(dbCtx, key); valueErr == nil { + fallbackValues[key] = value + } + } + enabled = strings.EqualFold(strings.TrimSpace(fallbackValues[openAIAdvancedSchedulerSettingKey]), "true") + stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true") + subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true") + lbTopKOverride = parsePositiveIntOverride(fallbackValues[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(fallbackValues) } } openAIAdvancedSchedulerSettingCache.Store(&cachedOpenAIAdvancedSchedulerSetting{ - enabled: enabled, - expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(), + enabled: enabled, + stickyWeightedEnabled: stickyWeightedEnabled, + subscriptionPriorityEnabled: subscriptionPriorityEnabled, + lbTopKOverride: lbTopKOverride, + weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(weightOverrides), + expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(), }) - return enabled, nil + return openAIAdvancedSchedulerRuntimeSettings{ + enabled: enabled, + stickyWeightedEnabled: stickyWeightedEnabled, + subscriptionPriorityEnabled: subscriptionPriorityEnabled, + lbTopKOverride: lbTopKOverride, + weightOverrides: weightOverrides, + }, nil }) - enabled, _ := result.(bool) - return enabled + settings, _ := result.(openAIAdvancedSchedulerRuntimeSettings) + return settings +} + +func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerEnabled(ctx context.Context) bool { + return s.openAIAdvancedSchedulerRuntimeSettings(ctx).enabled +} + +func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx context.Context) bool { + settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + return settings.enabled && settings.stickyWeightedEnabled +} + +func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerSubscriptionPriorityEnabled(ctx context.Context) bool { + settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + return settings.enabled && settings.subscriptionPriorityEnabled +} + +func openAIAdvancedSchedulerRuntimeSettingKeys() []string { + keys := []string{ + openAIAdvancedSchedulerSettingKey, + SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled, + SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled, + SettingKeyOpenAIAdvancedSchedulerLBTopK, + } + for _, spec := range openAIAdvancedSchedulerWeightOverrideSpecs() { + keys = append(keys, spec.key) + } + return keys +} + +type openAIAdvancedSchedulerWeightOverrideSpec struct { + key string + name string +} + +func openAIAdvancedSchedulerWeightOverrideSpecs() []openAIAdvancedSchedulerWeightOverrideSpec { + return []openAIAdvancedSchedulerWeightOverrideSpec{ + {key: SettingKeyOpenAIAdvancedSchedulerWeightPriority, name: "priority"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightLoad, name: "load"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightQueue, name: "queue"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightErrorRate, name: "error_rate"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightTTFT, name: "ttft"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightReset, name: "reset"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom, name: "quota_headroom"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse, name: "previous_response"}, + {key: SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky, name: "session_sticky"}, + } +} + +func parsePositiveIntOverride(raw string) int { + raw = strings.TrimSpace(raw) + if raw == "" { + return 0 + } + value, err := strconv.Atoi(raw) + if err != nil || value <= 0 { + return 0 + } + return value +} + +func parseOpenAIAdvancedSchedulerWeightOverrides(values map[string]string) map[string]float64 { + overrides := map[string]float64{} + for _, spec := range openAIAdvancedSchedulerWeightOverrideSpecs() { + raw := strings.TrimSpace(values[spec.key]) + if raw == "" { + continue + } + value, err := strconv.ParseFloat(raw, 64) + if err != nil || value < 0 || math.IsNaN(value) || math.IsInf(value, 0) { + continue + } + overrides[spec.name] = value + } + return overrides +} + +func cloneOpenAIAdvancedSchedulerWeightOverrides(in map[string]float64) map[string]float64 { + if len(in) == 0 { + return nil + } + out := make(map[string]float64, len(in)) + for key, value := range in { + out[key] = value + } + return out } func (s *OpenAIGatewayService) getOpenAIAccountScheduler(ctx context.Context) OpenAIAccountScheduler { @@ -1253,9 +1644,12 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler( requiredTransport OpenAIUpstreamTransport, requireCompact bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI) + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI, false) } +// SelectAccountWithSchedulerForCapability 按能力要求调度账号。 +// previousResponseCanMove 表示首包 input 可自行重建工具续链,previous_response_id 允许跨账号迁移 +// (粘性加权模式下改为加权偏好而非硬粘连)。 func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( ctx context.Context, groupID *int64, @@ -1266,13 +1660,14 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( requiredTransport OpenAIUpstreamTransport, requiredCapability OpenAIEndpointCapability, requireCompact bool, + previousResponseCanMove bool, platformOverride ...string, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { platform := PlatformOpenAI if len(platformOverride) > 0 { platform = platformOverride[0] } - return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform) + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform, previousResponseCanMove) } func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( @@ -1283,13 +1678,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( excludedIDs map[int64]struct{}, requiredCapability OpenAIImagesCapability, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false, PlatformOpenAI) + selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false, PlatformOpenAI, false) if err == nil && selection != nil && selection.Account != nil { return selection, decision, nil } // 如果要求 native 能力(如指定了模型)但没有可用的 APIKey 账号,回退到 basic(OAuth 账号) if requiredCapability == OpenAIImagesCapabilityNative { - return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false, PlatformOpenAI) + return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false, PlatformOpenAI, false) } return selection, decision, err } @@ -1306,6 +1701,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( requiredImageCapability OpenAIImagesCapability, requireCompact bool, platform string, + previousResponseCanMove bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { ctx = s.withOpenAIQuotaAutoPauseContext(ctx) platform = normalizeOpenAICompatiblePlatform(platform) @@ -1378,13 +1774,23 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( stickyAccountID = accountID } } + stickyWeighted := s.isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx) + subscriptionPriority := s.isOpenAIAdvancedSchedulerSubscriptionPriorityEnabled(ctx) + stickyPreviousAccountID := int64(0) + if stickyWeighted && previousResponseCanMove && strings.TrimSpace(previousResponseID) != "" && platform == PlatformOpenAI { + stickyPreviousAccountID = s.ResolveAccountIDByPreviousResponseIDForScheduler(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact) + } return scheduler.Select(ctx, OpenAIAccountScheduleRequest{ GroupID: groupID, Platform: platform, SessionHash: sessionHash, StickyAccountID: stickyAccountID, + StickyPreviousAccountID: stickyPreviousAccountID, + StickyWeighted: stickyWeighted, + SubscriptionPriority: subscriptionPriority, PreviousResponseID: previousResponseID, + PreviousResponseCanMove: previousResponseCanMove, RequestedModel: requestedModel, RequiredTransport: requiredTransport, RequiredCapability: requiredCapability, @@ -1473,6 +1879,20 @@ func (s *OpenAIGatewayService) openAIWSLBTopK() int { return 7 } +func (s *OpenAIGatewayService) openAIWSLBTopKForRequest(ctx context.Context) int { + base := s.openAIWSLBTopK() + settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // DB 覆盖值与 stickyWeighted/subscriptionPriority 一样受总开关门控: + // 关闭高级调度器后所有调用方(含管理页分数快照)都应回到配置/默认行为。 + if !settings.enabled { + return base + } + if settings.lbTopKOverride > 0 { + return settings.lbTopKOverride + } + return base +} + func (s *OpenAIGatewayService) openAIStickyEscapeConfig() openAIStickyEscapeConfig { if s != nil && s.cfg != nil { cfg := s.cfg.Gateway.OpenAIScheduler @@ -1514,6 +1934,8 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul TTFT: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT, Reset: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Reset, QuotaHeadroom: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom, + Previous: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse, + SessionSticky: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky, } } return GatewayOpenAIWSSchedulerScoreWeightsView{ @@ -1524,9 +1946,50 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul TTFT: 0.5, Reset: 0.0, QuotaHeadroom: 0.0, + Previous: 5.0, + SessionSticky: 3.0, } } +func (s *OpenAIGatewayService) openAIWSSchedulerWeightsForRequest(ctx context.Context) GatewayOpenAIWSSchedulerScoreWeightsView { + weights := s.openAIWSSchedulerWeights() + settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // 同 openAIWSLBTopKForRequest:总开关关闭时不应用 DB 覆盖值。 + if !settings.enabled { + return weights + } + return applyOpenAIAdvancedSchedulerWeightOverrides(weights, settings.weightOverrides) +} + +func applyOpenAIAdvancedSchedulerWeightOverrides( + weights GatewayOpenAIWSSchedulerScoreWeightsView, + overrides map[string]float64, +) GatewayOpenAIWSSchedulerScoreWeightsView { + for key, value := range overrides { + switch key { + case "priority": + weights.Priority = value + case "load": + weights.Load = value + case "queue": + weights.Queue = value + case "error_rate": + weights.ErrorRate = value + case "ttft": + weights.TTFT = value + case "reset": + weights.Reset = value + case "quota_headroom": + weights.QuotaHeadroom = value + case "previous_response": + weights.Previous = value + case "session_sticky": + weights.SessionSticky = value + } + } + return weights +} + type GatewayOpenAIWSSchedulerScoreWeightsView struct { Priority float64 Load float64 @@ -1536,6 +1999,149 @@ type GatewayOpenAIWSSchedulerScoreWeightsView struct { // Reset 倾向「会话窗口最早重置」的账号;0 表示关闭(默认)。 Reset float64 QuotaHeadroom float64 + Previous float64 + SessionSticky float64 +} + +type OpenAIAccountSchedulerScoreSnapshot struct { + BaseScore float64 + StickyScore float64 + StickyScoreInfinity bool + StickyWeightedEnabled bool +} + +func (s *RateLimitService) BuildOpenAIAccountSchedulerScoreSnapshot( + ctx context.Context, + accounts []*Account, + loadMap map[int64]*AccountLoadInfo, +) map[int64]OpenAIAccountSchedulerScoreSnapshot { + gateway := &OpenAIGatewayService{cfg: nil, rateLimitService: s} + if s != nil { + gateway.cfg = s.cfg + } + return buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, gateway.openAIWSSchedulerWeightsForRequest(ctx), gateway.isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx)) +} + +func BuildOpenAIAccountSchedulerScoreSnapshot( + accounts []*Account, + loadMap map[int64]*AccountLoadInfo, +) map[int64]OpenAIAccountSchedulerScoreSnapshot { + gateway := &OpenAIGatewayService{} + return buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, gateway.openAIWSSchedulerWeights(), false) +} + +func buildOpenAIAccountSchedulerScoreSnapshot( + accounts []*Account, + loadMap map[int64]*AccountLoadInfo, + weights GatewayOpenAIWSSchedulerScoreWeightsView, + stickyWeightedEnabled bool, +) map[int64]OpenAIAccountSchedulerScoreSnapshot { + if len(accounts) == 0 { + return nil + } + candidates := make([]openAIAccountCandidateScore, 0, len(accounts)) + for _, account := range accounts { + if account == nil { + continue + } + loadInfo := loadMap[account.ID] + if loadInfo == nil { + loadInfo = &AccountLoadInfo{AccountID: account.ID} + } + candidates = append(candidates, openAIAccountCandidateScore{ + account: account, + loadInfo: loadInfo, + errorRate: 0, + ttft: 0, + hasTTFT: false, + }) + } + if len(candidates) == 0 { + return nil + } + + minPriority, maxPriority := openAIAccountSchedulingPriority(candidates[0].account), openAIAccountSchedulingPriority(candidates[0].account) + maxWaiting := 1 + for i := range candidates { + candidate := &candidates[i] + candidate.priority = openAIAccountSchedulingPriority(candidate.account) + if candidate.priority < minPriority { + minPriority = candidate.priority + } + if candidate.priority > maxPriority { + maxPriority = candidate.priority + } + if candidate.loadInfo.WaitingCount > maxWaiting { + maxWaiting = candidate.loadInfo.WaitingCount + } + } + + minResetRemaining, maxResetRemaining := 0.0, 0.0 + hasResetSample := false + now := time.Now() + if weights.Reset > 0 { + for _, candidate := range candidates { + end := candidate.account.SessionWindowEnd + if end == nil || !now.Before(*end) { + continue + } + remaining := end.Sub(now).Seconds() + if !hasResetSample { + minResetRemaining, maxResetRemaining = remaining, remaining + hasResetSample = true + continue + } + if remaining < minResetRemaining { + minResetRemaining = remaining + } + if remaining > maxResetRemaining { + maxResetRemaining = remaining + } + } + } + + result := make(map[int64]OpenAIAccountSchedulerScoreSnapshot, len(candidates)) + for _, candidate := range candidates { + priorityFactor := 1.0 + if maxPriority > minPriority { + priorityFactor = 1 - float64(candidate.priority-minPriority)/float64(maxPriority-minPriority) + } + loadFactor := 1 - clamp01(float64(candidate.loadInfo.LoadRate)/100.0) + queueFactor := 1 - clamp01(float64(candidate.loadInfo.WaitingCount)/float64(maxWaiting)) + errorFactor := 1.0 + ttftFactor := 0.5 + resetFactor := 0.0 + if weights.Reset > 0 && hasResetSample { + if end := candidate.account.SessionWindowEnd; end != nil && now.Before(*end) { + if maxResetRemaining > minResetRemaining { + resetFactor = 1 - clamp01((end.Sub(now).Seconds()-minResetRemaining)/(maxResetRemaining-minResetRemaining)) + } else { + resetFactor = 1 + } + } + } + quotaHeadroomFactor := 0.0 + if weights.QuotaHeadroom > 0 { + quotaHeadroomFactor = openAIQuotaHeadroomFactor(candidate.account, now) + } + baseScore := weights.Priority*priorityFactor + + weights.Load*loadFactor + + weights.Queue*queueFactor + + weights.ErrorRate*errorFactor + + weights.TTFT*ttftFactor + + weights.Reset*resetFactor + + weights.QuotaHeadroom*quotaHeadroomFactor + score := OpenAIAccountSchedulerScoreSnapshot{ + BaseScore: baseScore, + StickyWeightedEnabled: stickyWeightedEnabled, + StickyScoreInfinity: !stickyWeightedEnabled, + } + if stickyWeightedEnabled { + score.StickyScore = baseScore + weights.Previous + weights.SessionSticky + } + result[candidate.account.ID] = score + } + return result } func openAIQuotaHeadroomFactor(account *Account, now time.Time) float64 { diff --git a/backend/internal/service/openai_account_scheduler_reset_test.go b/backend/internal/service/openai_account_scheduler_reset_test.go index 50155f20d6..f1270c8caa 100644 --- a/backend/internal/service/openai_account_scheduler_reset_test.go +++ b/backend/internal/service/openai_account_scheduler_reset_test.go @@ -1,6 +1,7 @@ package service import ( + "context" "testing" "time" @@ -49,7 +50,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightPrefersSoonestReset(t *testing.T) } sched := openAIResetTestScheduler(5.0) - plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) + plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) scores := openAIPlanScores(plan) require.Greater(t, scores[2], scores[1], "重置时间最早的账号(ID=2)得分更高") } @@ -65,7 +66,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightZeroNoEffect(t *testing.T) { } sched := openAIResetTestScheduler(0.0) - plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) + plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) scores := openAIPlanScores(plan) require.Equal(t, scores[1], scores[2], "Reset 权重为 0 时两账号得分相同") } @@ -80,7 +81,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightIgnoresNilWindow(t *testing.T) { } sched := openAIResetTestScheduler(5.0) - plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) + plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) scores := openAIPlanScores(plan) require.Greater(t, scores[2], scores[1], "拥有活跃窗口的账号得分高于无窗口账号") } @@ -161,7 +162,7 @@ func TestBuildOpenAIAccountLoadPlan_QuotaHeadroomPrefersHigher7dRemaining(t *tes } sched := openAIQuotaHeadroomTestScheduler(1.0) - plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) + plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) scores := openAIPlanScores(plan) require.Greater(t, scores[2], scores[1], "7d 剩余额度更高的账号得分应更高") } @@ -190,7 +191,7 @@ func TestBuildOpenAIAccountLoadPlan_QuotaHeadroomZeroNoEffect(t *testing.T) { } sched := openAIResetTestScheduler(0) - plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) + plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{}) scores := openAIPlanScores(plan) require.Equal(t, scores[1], scores[2], "quota_headroom 权重为 0 时不应影响打分") } diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 0255dbbde3..2ff7c25e5d 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -183,6 +183,17 @@ func newSchedulerTestOpenAIWSV2Config() *config.Config { return cfg } +func newSchedulerTestSubscriptionPriorityConfig() *config.Config { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0 + return cfg +} + type openAIAdvancedSchedulerSettingRepoStub struct { values map[string]string } @@ -210,8 +221,14 @@ func (s *openAIAdvancedSchedulerSettingRepoStub) Set(context.Context, string, st panic("unexpected call to Set") } -func (s *openAIAdvancedSchedulerSettingRepoStub) GetMultiple(context.Context, []string) (map[string]string, error) { - panic("unexpected call to GetMultiple") +func (s *openAIAdvancedSchedulerSettingRepoStub) GetMultiple(_ context.Context, keys []string) (map[string]string, error) { + result := make(map[string]string, len(keys)) + for _, key := range keys { + if value, err := s.GetValue(context.Background(), key); err == nil { + result[key] = value + } + } + return result, nil } func (s *openAIAdvancedSchedulerSettingRepoStub) SetMultiple(context.Context, map[string]string) error { @@ -226,7 +243,7 @@ func (s *openAIAdvancedSchedulerSettingRepoStub) Delete(context.Context, string) panic("unexpected call to Delete") } -func newOpenAIAdvancedSchedulerRateLimitService(enabled string) *RateLimitService { +func newOpenAIAdvancedSchedulerRateLimitService(enabled string, values ...string) *RateLimitService { resetOpenAIAdvancedSchedulerSettingCacheForTest() repo := &openAIAdvancedSchedulerSettingRepoStub{ values: map[string]string{}, @@ -234,6 +251,12 @@ func newOpenAIAdvancedSchedulerRateLimitService(enabled string) *RateLimitServic if enabled != "" { repo.values[openAIAdvancedSchedulerSettingKey] = enabled } + if len(values) > 0 && values[0] != "" { + repo.values[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] = values[0] + } + if len(values) > 1 && values[1] != "" { + repo.values[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] = values[1] + } return &RateLimitService{ settingService: NewSettingService(repo, &config.Config{}), } @@ -266,6 +289,45 @@ func (s *openAISnapshotCacheStub) GetAccount(ctx context.Context, accountID int6 return &cloned, nil } +func TestOpenAIGatewayService_OpenAIAdvancedSchedulerRuntimeSettings_DBOverridesConfig(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + defer resetOpenAIAdvancedSchedulerSettingCacheForTest() + + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 11 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights = config.GatewayOpenAIWSSchedulerScoreWeights{ + Priority: 1, + Load: 2, + Queue: 3, + ErrorRate: 4, + TTFT: 5, + Reset: 6, + QuotaHeadroom: 7, + PreviousResponse: 8, + SessionSticky: 9, + } + repo := &openAIAdvancedSchedulerSettingRepoStub{ + values: map[string]string{ + openAIAdvancedSchedulerSettingKey: "true", + SettingKeyOpenAIAdvancedSchedulerLBTopK: "3", + SettingKeyOpenAIAdvancedSchedulerWeightPriority: "2.5", + SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: "12", + }, + } + svc := &OpenAIGatewayService{ + cfg: cfg, + rateLimitService: &RateLimitService{settingService: NewSettingService(repo, cfg)}, + } + + ctx := context.Background() + require.Equal(t, 3, svc.openAIWSLBTopKForRequest(ctx)) + weights := svc.openAIWSSchedulerWeightsForRequest(ctx) + require.Equal(t, 2.5, weights.Priority) + require.Equal(t, 2.0, weights.Load) + require.Equal(t, 12.0, weights.Previous) + require.Equal(t, 9.0, weights.SessionSticky) +} + func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabledUsesLegacyLoadAwareness(t *testing.T) { resetOpenAIAdvancedSchedulerSettingCacheForTest() @@ -467,6 +529,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -510,6 +573,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + false, PlatformGrok, ) require.NoError(t, err) @@ -584,6 +648,248 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPrev require.True(t, decision.StickyPreviousHit) } +func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedSessionInTopKUsesStickyFirst(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101071) + accounts := []Account{ + { + ID: 37101, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 100, + GroupIDs: []int64{groupID}, + }, + { + ID: 37102, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + }, + } + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 2 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3 + cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{ + "openai:session_hash_weighted_topk": 37101, + }} + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: cache, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_hash_weighted_topk", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37101), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.True(t, decision.StickySessionHit) + require.Equal(t, 2, decision.TopK) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousRequiresMovableContext(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101072) + accounts := []Account{ + { + ID: 37111, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 100, + GroupIDs: []int64{groupID}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + }, + }, + { + ID: 37112, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + }, + }, + } + cfg := newSchedulerTestOpenAIWSV2Config() + cfg.Gateway.OpenAIWS.LBTopK = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + store := svc.getOpenAIWSStateStore() + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_weighted_unmovable", 37111, time.Hour)) + + selection, decision, err := svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "resp_weighted_unmovable", + "", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + OpenAIEndpointCapabilityChatCompletions, + false, + false, + PlatformOpenAI, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37111), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerPreviousResponse, decision.Layer) + require.True(t, decision.StickyPreviousHit) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + + selection, decision, err = svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "resp_weighted_unmovable", + "", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + OpenAIEndpointCapabilityChatCompletions, + false, + true, + PlatformOpenAI, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37112), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.False(t, decision.StickyPreviousHit) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_PreviousResponseCompactUnsupportedDeletesBinding(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101073) + accounts := []Account{ + { + ID: 37121, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + "openai_compact_mode": OpenAICompactModeForceOff, + }, + }, + { + ID: 37122, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + "openai_compact_mode": OpenAICompactModeForceOn, + }, + }, + } + cfg := newSchedulerTestOpenAIWSV2Config() + cfg.Gateway.OpenAIWS.LBTopK = 2 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + store := svc.getOpenAIWSStateStore() + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_compact_unsupported", 37121, time.Hour)) + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "resp_compact_unsupported", + "", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + true, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37122), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.False(t, decision.StickyPreviousHit) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + + accountID, err := store.GetResponseAccount(ctx, groupID, "resp_compact_unsupported") + require.NoError(t, err) + require.Zero(t, accountID) +} + func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) { resetOpenAIAdvancedSchedulerSettingCacheForTest() @@ -635,6 +941,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -708,6 +1015,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -1560,6 +1868,217 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeDisa require.True(t, decision.StickySessionHit) } +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityChoosesSubscriptionPoolFirst(t *testing.T) { + ctx := context.Background() + groupID := int64(10120) + accounts := []Account{ + { + ID: 21601, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "plus"}, + }, + { + ID: 21602, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{21601: true, 21602: true}, + loadMap: map[int64]*AccountLoadInfo{ + 21601: {AccountID: 21601, LoadRate: 90, WaitingCount: 1}, + 21602: {AccountID: 21602, LoadRate: 0, WaitingCount: 0}, + }, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_first", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(21601), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.Equal(t, 1, decision.TopK) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityFallsBackWhenSubscriptionFull(t *testing.T) { + ctx := context.Background() + groupID := int64(10121) + accounts := []Account{ + { + ID: 21611, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "team"}, + }, + { + ID: 21612, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 9, + GroupIDs: []int64{groupID}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{21611: false, 21612: true}, + loadMap: map[int64]*AccountLoadInfo{ + 21611: {AccountID: 21611, LoadRate: 0, WaitingCount: 0}, + 21612: {AccountID: 21612, LoadRate: 90, WaitingCount: 1}, + }, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_fallback", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(21612), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.True(t, selection.Acquired) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityDisabledUsesScore(t *testing.T) { + ctx := context.Background() + groupID := int64(10122) + accounts := []Account{ + { + ID: 21621, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "pro"}, + }, + { + ID: 21622, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{21621: true, 21622: true}, + loadMap: map[int64]*AccountLoadInfo{ + 21621: {AccountID: 21621, LoadRate: 90, WaitingCount: 1}, + 21622: {AccountID: 21622, LoadRate: 0, WaitingCount: 0}, + }, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "false"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_disabled", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(21622), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_UsesAccountPriorityWithinGroupPool(t *testing.T) { + ctx := context.Background() + groupID := int64(10123) + accounts := []Account{ + { + ID: 21631, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + AccountGroups: []AccountGroup{ + {AccountID: 21631, GroupID: groupID, Priority: 100}, + }, + GroupIDs: []int64{groupID}, + }, + { + ID: 21632, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 100000, + AccountGroups: []AccountGroup{ + {AccountID: 21632, GroupID: groupID, Priority: 1}, + }, + GroupIDs: []int64{groupID}, + }, + } + cfg := newSchedulerTestSubscriptionPriorityConfig() + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 0 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0 + svc := &OpenAIGatewayService{ + accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_group_priority", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(21631), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + func TestDefaultOpenAIAccountScheduler_ShouldEscapeStickyAccount_ThresholdBoundary(t *testing.T) { stats := newOpenAIAccountRuntimeStats() accountID := int64(21501) @@ -2421,3 +2940,141 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t * func int64PtrForTest(v int64) *int64 { return &v } + +func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedFallbackSkipsOutOfGroupStickyAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101081) + otherGroupID := int64(101082) + accounts := []Account{ + { + ID: 38001, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + }, + { + // 会话粘连绑定指向的账号已被移出请求分组(绑定 TTL 内账号改组的场景)。 + ID: 38002, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{otherGroupID}, + }, + } + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 2 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3 + cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{ + "openai:session_weighted_out_of_group": 38002, + }} + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38001: false, 38002: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}}, + cache: cache, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_weighted_out_of_group", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + // 组内唯一候选 38001 满并发:必须返回其等待计划,绝不能把请求泄漏到组外的粘连账号 38002。 + require.Equal(t, int64(38001), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38001), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + // 失效的粘连绑定应被清理,避免后续请求反复走同一条泄漏路径。 + require.Positive(t, cache.deletedSessions["openai:session_weighted_out_of_group"]) +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityWaitsOnBusySubscriptionWhenRegularUnusable(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101091) + accounts := []Account{ + { + // 订阅账号:支持 compact,但并发已满(busy-but-waitable)。 + ID: 38011, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "team"}, + Extra: map[string]any{"openai_compact_supported": true}, + }, + { + // 常规账号:明确不支持 compact,无法服务本次请求。 + ID: 38012, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 9, + GroupIDs: []int64{groupID}, + Extra: map[string]any{"openai_compact_supported": false}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38011: false, 38012: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_subscription_wait", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + true, + ) + // 常规池无可用候选时,忙碌的订阅账号应产生等待计划,而不是直接返回 no available accounts。 + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(38011), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38011), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} diff --git a/backend/internal/service/openai_apikey_responses_probe.go b/backend/internal/service/openai_apikey_responses_probe.go index 64f963ab9b..10cf050029 100644 --- a/backend/internal/service/openai_apikey_responses_probe.go +++ b/backend/internal/service/openai_apikey_responses_probe.go @@ -149,6 +149,9 @@ func (s *AccountTestService) ProbeOpenAIAPIKeyResponsesSupport(ctx context.Conte req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Accept", "application/json") + // 账号级请求头覆写:能力探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_client_restriction_detector.go b/backend/internal/service/openai_client_restriction_detector.go index abca88ce66..8a8097c879 100644 --- a/backend/internal/service/openai_client_restriction_detector.go +++ b/backend/internal/service/openai_client_restriction_detector.go @@ -1,6 +1,7 @@ package service import ( + "fmt" "net/http" "github.com/Wei-Shaw/sub2api/internal/config" @@ -8,6 +9,11 @@ import ( "github.com/gin-gonic/gin" ) +// CodexOfficialClientsOnlyMessage 是 codex_cli_only 拒绝时面向客户端的通用兜底文案。 +// 仅当拒绝原因不是「可解析版本但越界」(VersionTooLow/VersionTooHigh)时使用: +// 未命中官方/黑名单/缺指纹/版本无法识别都沿用这句(避免向伪装客户端泄露门控细节)。 +const CodexOfficialClientsOnlyMessage = "This account only allows Codex official clients" + const ( // CodexClientRestrictionReasonDisabled 表示账号未开启 codex_cli_only。 CodexClientRestrictionReasonDisabled = "codex_cli_only_disabled" @@ -51,6 +57,13 @@ type CodexClientRestrictionDetectionResult struct { Enabled bool Matched bool Reason string + // DetectedVersion 是从官方 UA 解析出的 Codex 引擎版本;仅在版本门拒绝 + // (VersionTooLow / VersionTooHigh) 时填充,供面向客户端的差异化文案使用。 + DetectedVersion string + // MinCodexVersion 是触发 VersionTooLow 时的最低要求版本(来自策略快照)。 + MinCodexVersion string + // MaxCodexVersion 是触发 VersionTooHigh 时的最高允许版本(来自策略快照)。 + MaxCodexVersion string } // CodexClientRestrictionDetector 定义 codex_cli_only 统一检测入口。 @@ -127,10 +140,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionUndetectable} } if policy.MinCodexVersion != "" && CompareVersions(ver, policy.MinCodexVersion) < 0 { - return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooLow} + return CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: false, + Reason: CodexClientRestrictionReasonVersionTooLow, + DetectedVersion: ver, + MinCodexVersion: policy.MinCodexVersion, + } } if policy.MaxCodexVersion != "" && CompareVersions(ver, policy.MaxCodexVersion) > 0 { - return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooHigh} + return CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: false, + Reason: CodexClientRestrictionReasonVersionTooHigh, + DetectedVersion: ver, + MaxCodexVersion: policy.MaxCodexVersion, + } } } @@ -145,3 +170,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A return CodexClientRestrictionDetectionResult{Enabled: true, Matched: true, Reason: reason} } + +// CodexClientRestrictionMessage 把检测结果映射为面向客户端的 403 文案。 +// 仅版本越界(VersionTooLow/VersionTooHigh)给出带实际版本号与边界的差异化提示—— +// 这类请求其实已被识别为官方 Codex(命中官方 UA/originator),再回「只允许官方客户端」会误导; +// 其余拒绝原因统一沿用通用兜底句,不暴露门控细节。 +func CodexClientRestrictionMessage(r CodexClientRestrictionDetectionResult) string { + switch r.Reason { + case CodexClientRestrictionReasonVersionTooLow: + return fmt.Sprintf( + "Your Codex version (%s) is below the minimum required version (%s). Please update Codex.", + r.DetectedVersion, r.MinCodexVersion) + case CodexClientRestrictionReasonVersionTooHigh: + return fmt.Sprintf( + "Your Codex version (%s) exceeds the maximum allowed version (%s). Please downgrade Codex to %s or lower.", + r.DetectedVersion, r.MaxCodexVersion, r.MaxCodexVersion) + default: + return CodexOfficialClientsOnlyMessage + } +} diff --git a/backend/internal/service/openai_client_restriction_detector_test.go b/backend/internal/service/openai_client_restriction_detector_test.go index 291c79f6bf..6c79432ae2 100644 --- a/backend/internal/service/openai_client_restriction_detector_test.go +++ b/backend/internal/service/openai_client_restriction_detector_test.go @@ -284,6 +284,66 @@ func TestDetect_V3_AppServerAndSkipAndVersionScope(t *testing.T) { }) } +func TestDetect_VersionGateCarriesVersionFields(t *testing.T) { + gin.SetMode(gin.TestMode) + d := NewOpenAICodexClientRestrictionDetector(nil) + acc := func() *Account { + return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}} + } + + t.Run("版本太低:携带 DetectedVersion + MinCodexVersion", func(t *testing.T) { + c := newCodexDetectorTestContext("codex_cli_rs/0.39.0 (x)", "") + r := d.Detect(c, acc(), CodexRestrictionPolicy{MinCodexVersion: "0.42.0"}, nil) + require.False(t, r.Matched) + require.Equal(t, CodexClientRestrictionReasonVersionTooLow, r.Reason) + require.Equal(t, "0.39.0", r.DetectedVersion) + require.Equal(t, "0.42.0", r.MinCodexVersion) + }) + + t.Run("版本太高:携带 DetectedVersion + MaxCodexVersion", func(t *testing.T) { + c := newCodexDetectorTestContext("codex_cli_rs/0.45.0 (x)", "") + r := d.Detect(c, acc(), CodexRestrictionPolicy{MaxCodexVersion: "0.42.0"}, nil) + require.False(t, r.Matched) + require.Equal(t, CodexClientRestrictionReasonVersionTooHigh, r.Reason) + require.Equal(t, "0.45.0", r.DetectedVersion) + require.Equal(t, "0.42.0", r.MaxCodexVersion) + }) +} + +func TestCodexClientRestrictionMessage(t *testing.T) { + t.Run("版本太低:带实际版本与最低要求", func(t *testing.T) { + msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{ + Reason: CodexClientRestrictionReasonVersionTooLow, + DetectedVersion: "0.39.0", + MinCodexVersion: "0.42.0", + }) + require.Equal(t, "Your Codex version (0.39.0) is below the minimum required version (0.42.0). Please update Codex.", msg) + }) + + t.Run("版本太高:带实际版本与最高允许", func(t *testing.T) { + msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{ + Reason: CodexClientRestrictionReasonVersionTooHigh, + DetectedVersion: "0.45.0", + MaxCodexVersion: "0.42.0", + }) + require.Equal(t, "Your Codex version (0.45.0) exceeds the maximum allowed version (0.42.0). Please downgrade Codex to 0.42.0 or lower.", msg) + }) + + t.Run("无法识别版本:保持原通用句", func(t *testing.T) { + msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{ + Reason: CodexClientRestrictionReasonVersionUndetectable, + }) + require.Equal(t, "This account only allows Codex official clients", msg) + }) + + t.Run("未命中官方:保持原通用句", func(t *testing.T) { + msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{ + Reason: CodexClientRestrictionReasonNotMatchedUA, + }) + require.Equal(t, "This account only allows Codex official clients", msg) + }) +} + func TestDetect_EngineFingerprintSignals(t *testing.T) { gin.SetMode(gin.TestMode) det := NewOpenAICodexClientRestrictionDetector(&config.Config{}) diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 0ece5e44ff..0666293deb 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -9,6 +9,9 @@ import ( ) var codexModelMap = map[string]string{ + "gpt-5.6-sol": "gpt-5.6-sol", + "gpt-5.6-terra": "gpt-5.6-terra", + "gpt-5.6-luna": "gpt-5.6-luna", "gpt-5.5": "gpt-5.5", "gpt-5.5-pro": "gpt-5.5-pro", "codex-auto-review": "codex-auto-review", @@ -54,6 +57,9 @@ var codexVersionModelPrefixes = []struct { prefix string target string }{ + {prefix: "gpt-5.6-sol", target: "gpt-5.6-sol"}, + {prefix: "gpt-5.6-terra", target: "gpt-5.6-terra"}, + {prefix: "gpt-5.6-luna", target: "gpt-5.6-luna"}, {prefix: "gpt-5.3-codex-spark", target: "gpt-5.3-codex-spark"}, {prefix: "gpt-5.3-codex", target: "gpt-5.3-codex"}, {prefix: "gpt-5.4-mini", target: "gpt-5.4-mini"}, @@ -607,18 +613,21 @@ func hasOpenAIImageGenerationTool(reqBody map[string]any) bool { return false } -// stripCodexSparkImageGenerationTools removes image_generation tool entries from -// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400 -// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so -// it must be dropped for spark. When the tools list becomes empty the key is removed. -// Returns true when the body was modified. -func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool { +func stripOpenAIImageGenerationTools(reqBody map[string]any) bool { rawTools, ok := reqBody["tools"] if !ok || rawTools == nil { + if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + delete(reqBody, "tool_choice") + return true + } return false } tools, ok := rawTools.([]any) if !ok { + if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + delete(reqBody, "tool_choice") + return true + } return false } filtered := make([]any, 0, len(tools)) @@ -631,17 +640,31 @@ func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool { } filtered = append(filtered, rawTool) } - if !removed { + if !removed && !openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { return false } - if len(filtered) == 0 { - delete(reqBody, "tools") - } else { - reqBody["tools"] = filtered + if removed { + if len(filtered) == 0 { + delete(reqBody, "tools") + } else { + reqBody["tools"] = filtered + } + } + if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + delete(reqBody, "tool_choice") } return true } +// stripCodexSparkImageGenerationTools removes image_generation tool entries from +// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400 +// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so +// it must be dropped for spark. When the tools list becomes empty the key is removed. +// Returns true when the body was modified. +func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool { + return stripOpenAIImageGenerationTools(reqBody) +} + func hasOpenAIInputImage(reqBody map[string]any) bool { if reqBody == nil { return false diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go index 0fb3fff1f7..fb2dc5ccbb 100644 --- a/backend/internal/service/openai_embeddings.go +++ b/backend/internal/service/openai_embeddings.go @@ -82,6 +82,9 @@ func (s *OpenAIGatewayService) ForwardEmbeddings( upstreamReq.Header.Set("user-agent", customUA) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + proxyURL := "" if account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 6bcb6718b7..348213a992 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -166,6 +166,9 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( upstreamReq.Header.Set("user-agent", "sub2api-grok/1.0") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + // 6. Send request proxyURL := "" if account.Proxy != nil { diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index 4a01b143e9..7518a6073a 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -231,6 +231,9 @@ func (s *OpenAIGatewayService) buildInputTokensUpstreamRequest( } } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index a7004f5d54..697d89e81c 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -1306,6 +1306,66 @@ func TestOpenAIGatewayServiceRecordUsage_ChannelMappedOverridesBillingModelWhenM require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero") } +func TestOpenAIGatewayServiceRecordUsage_ResponsesMappedBillingModelHonorsBillingModelSource(t *testing.T) { + usage := OpenAIUsage{InputTokens: 20, OutputTokens: 10} + tokens := UsageTokens{InputTokens: 20, OutputTokens: 10} + + tests := []struct { + name string + billingModelSource string + wantBillingModel string + }{ + { + name: "upstream uses mapped billing model", + billingModelSource: BillingModelSourceUpstream, + wantBillingModel: "gpt-5.5", + }, + { + name: "requested overrides mapped billing model", + billingModelSource: BillingModelSourceRequested, + wantBillingModel: "gpt-5.4", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + + expectedCost, err := svc.billingService.CalculateCost(tt.wantBillingModel, tokens, 1.1) + require.NoError(t, err) + + err = svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_mapped_billing_model_source", + Model: "gpt-5.4", + BillingModel: "gpt-5.5", + UpstreamModel: "gpt-5.5", + Usage: usage, + Duration: time.Second, + }, + APIKey: &APIKey{ID: 10}, + User: &User{ID: 20}, + Account: &Account{ID: 30}, + ChannelUsageFields: ChannelUsageFields{ + OriginalModel: "gpt-5.4", + ChannelMappedModel: "gpt-5.4", + BillingModelSource: tt.billingModelSource, + }, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, "gpt-5.4", usageRepo.lastLog.Model) + require.InDelta(t, expectedCost.ActualCost, usageRepo.lastLog.ActualCost, 1e-12) + require.InDelta(t, expectedCost.ActualCost, userRepo.lastAmount, 1e-12) + require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero") + }) + } +} + func TestOpenAIGatewayServiceRecordUsage_BillsCompactOpenAIModelAlias(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} @@ -1743,6 +1803,52 @@ func TestOpenAIGatewayServiceRecordUsage_ImageIndependentMultiplierUsesImageRate require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode) } +func TestGrokVideoMediaBillingUsesImageRateMultiplier(t *testing.T) { + mediaPrice2K := 0.4 + groupID := int64(126) + + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "video-request-123", + ResponseID: "video-request-123", + Model: "grok-imagine-video-1.5", + BillingModel: "grok-imagine-video-1.5", + // The usage schema has no separate video count; video generation is billed as one media unit. + ImageCount: 1, + ImageSize: ImageBillingSize2K, + Duration: time.Second, + }, + APIKey: &APIKey{ + ID: 10126, + GroupID: i64p(groupID), + Group: &Group{ + ID: groupID, + Platform: PlatformGrok, + RateMultiplier: 0.15, + ImageRateIndependent: true, + ImageRateMultiplier: 0.5, + ImagePrice2K: &mediaPrice2K, + }, + }, + User: &User{ID: 20126}, + Account: &Account{ID: 30126, Platform: PlatformGrok}, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, "grok-imagine-video-1.5", usageRepo.lastLog.Model) + require.Equal(t, 1, usageRepo.lastLog.ImageCount) + require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize) + require.InDelta(t, 0.4, usageRepo.lastLog.TotalCost, 1e-12) + require.InDelta(t, 0.2, usageRepo.lastLog.ActualCost, 1e-12) + require.InDelta(t, 0.5, usageRepo.lastLog.RateMultiplier, 1e-12) + require.NotNil(t, usageRepo.lastLog.BillingMode) + require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode) +} + func TestOpenAIGatewayServiceRecordUsage_ChannelImageBillingUsesImageCountAndSharedMultiplier(t *testing.T) { groupID := int64(123) usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index d33df4c19d..c499bec778 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -138,6 +138,9 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( upstreamReq.Header.Set("user-agent", customUA) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + proxyURL := "" if account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index cd84790b98..645b31992a 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -2617,7 +2617,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco c.JSON(http.StatusForbidden, gin.H{ "error": gin.H{ "type": "forbidden_error", - "message": "This account only allows Codex official clients", + "message": CodexClientRestrictionMessage(restrictionResult), }, }) return nil, errors.New("codex_cli_only restriction: only codex official clients are allowed") @@ -2738,8 +2738,25 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if apiKey != nil { imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group) } - codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) - imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) + codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow + if isCodexCLI { + codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() + } + codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) + var imageIntent bool + if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if stripOpenAIImageGenerationTools(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Stripped /responses image_generation tool for Codex client by account policy") + } + imageIntent = IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, decoded) + } else { + imageIntent = IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) + } if imageIntent && !imageGenerationAllowed { MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}}) @@ -3216,6 +3233,9 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco wsAttempts, ) wsResult.UpstreamModel = upstreamModel + if wsResult.BillingModel == "" { + wsResult.BillingModel = billingModel + } if wsResult.ImageCount > 0 { wsResult.ImageSize = imageSizeTier wsResult.ImageInputSize = imageInputSize @@ -3363,6 +3383,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco ResponseID: responseID, Usage: *usage, Model: originalModel, + BillingModel: billingModel, UpstreamModel: upstreamModel, ServiceTier: serviceTier, ReasoningEffort: reasoningEffort, @@ -3762,6 +3783,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( req.Header.Set("content-type", "application/json") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } @@ -4547,6 +4571,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. req.Header.Set("content-type", "application/json") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } diff --git a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go index 23a1750021..7eb133c125 100644 --- a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go +++ b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go @@ -59,6 +59,52 @@ func TestOpenAIGatewayService_GetCodexClientRestrictionDetector(t *testing.T) { }) } +func TestOpenAIGatewayService_Forward_VersionGateMessage(t *testing.T) { + gin.SetMode(gin.TestMode) + + newCtx := func() (*httptest.ResponseRecorder, *gin.Context) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + return rec, c + } + account := func() *Account { + return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}} + } + body := []byte(`{"model":"gpt-5.1-codex"}`) + + t.Run("版本太低:返回带版本号的差异化文案", func(t *testing.T) { + rec, c := newCtx() + svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: false, + Reason: CodexClientRestrictionReasonVersionTooLow, + DetectedVersion: "0.39.0", + MinCodexVersion: "0.42.0", + }}} + + _, err := svc.Forward(context.Background(), c, account(), body) + require.Error(t, err) + require.Equal(t, http.StatusForbidden, rec.Code) + require.Contains(t, rec.Body.String(), "Your Codex version (0.39.0) is below the minimum required version (0.42.0)") + require.NotContains(t, rec.Body.String(), "This account only allows Codex official clients") + }) + + t.Run("未命中官方:仍返回通用兜底文案", func(t *testing.T) { + rec, c := newCtx() + svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: false, + Reason: CodexClientRestrictionReasonNotMatchedUA, + }}} + + _, err := svc.Forward(context.Background(), c, account(), body) + require.Error(t, err) + require.Equal(t, http.StatusForbidden, rec.Code) + require.Contains(t, rec.Body.String(), "This account only allows Codex official clients") + }) +} + func TestGetAPIKeyIDFromContext(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index aee69fffa9..1dde60c9f0 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -196,6 +196,144 @@ func TestOpenAIGatewayService_Forward_MappedImageModelUsesImageGate(t *testing.T require.Equal(t, http.StatusForbidden, rec.Code) } +func TestOpenAIGatewayService_Forward_TextResponsesSetsBillingModelToMappedModel(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_text_mapped_billing"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"resp_text_mapped","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`, + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 4, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + "model_mapping": map[string]any{"gpt-5.4": "gpt-5.5"}, + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gpt-5.4", result.Model) + require.Equal(t, "gpt-5.5", result.BillingModel) + require.Equal(t, "gpt-5.5", result.UpstreamModel) + require.Equal(t, "gpt-5.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, 0, result.ImageCount) +} + +func TestOpenAIGatewayService_Forward_TextResponsesWithoutMappingKeepsRequestedBillingModel(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_text_unmapped_billing"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_text_unmapped","object":"response","model":"gpt-5.4","status":"completed","usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 4, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`)) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gpt-5.4", result.Model) + require.Equal(t, "gpt-5.4", result.BillingModel) + require.Equal(t, "gpt-5.4", result.UpstreamModel) +} + +func TestOpenAIGatewayService_Forward_TextResponsesBillingModelMatchesChatCompletions(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + account := &Account{ + ID: 5, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + "model_mapping": map[string]any{"gpt-5.4": "gpt-5.5"}, + }, + Extra: map[string]any{"use_responses_api": true}, + } + + responsesUpstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_responses_mapped_billing"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"resp_native","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`, + )), + }, + } + responsesSvc := &OpenAIGatewayService{cfg: cfg, httpUpstream: responsesUpstream} + responsesRecorder := httptest.NewRecorder() + responsesCtx, _ := gin.CreateTestContext(responsesRecorder) + responsesCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(responsesCtx, OpenAIClientTransportHTTP) + responsesResult, err := responsesSvc.Forward(context.Background(), responsesCtx, account, []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`)) + require.NoError(t, err) + require.NotNil(t, responsesResult) + + chatUpstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_chat_mapped_billing"}}, + Body: io.NopCloser(strings.NewReader( + `data: {"type":"response.completed","response":{"id":"resp_chat","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}}` + "\n\n", + )), + }, + } + chatSvc := &OpenAIGatewayService{cfg: cfg, httpUpstream: chatUpstream} + chatRecorder := httptest.NewRecorder() + chatCtx, _ := gin.CreateTestContext(chatRecorder) + chatCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/chat/completions", nil) + chatResult, err := chatSvc.ForwardAsChatCompletions(context.Background(), chatCtx, account, []byte(`{"model":"gpt-5.4","stream":false,"messages":[{"role":"user","content":"hello"}]}`), "", "") + require.NoError(t, err) + require.NotNil(t, chatResult) + + require.Equal(t, chatResult.BillingModel, responsesResult.BillingModel) + require.Equal(t, "gpt-5.5", responsesResult.BillingModel) + require.Equal(t, "gpt-5.5", chatResult.BillingModel) +} + func TestOpenAIGatewayService_Forward_TextDataImageDoesNotForceMapMarshal(t *testing.T) { gin.SetMode(gin.TestMode) upstream := &httpUpstreamRecorder{ diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index 6061a0bf9c..31edd36097 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -152,6 +152,45 @@ func TestOpenAIGatewayServiceForward_ExplicitImageToolWorksWithBridgeDisabled(t require.NotContains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_AccountPolicyStripsExplicitImageTool(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_stripped_image","model":"gpt-5.4","usage":{"input_tokens":2,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.98.0") + account := newOpenAIImageGenerationControlTestAccount() + account.Extra = map[string]any{ + featureKeyCodexImageGenerationExplicitToolPolicy: codexImageGenerationExplicitToolPolicyStrip, + } + body := []byte(`{ + "model":"gpt-5.4", + "input":"draw", + "stream":false, + "tools":[ + {"type":"function","name":"shell","parameters":{"type":"object"}}, + {"type":"image_generation","format":"jpeg"} + ], + "tool_choice":{"type":"image_generation"} + }`) + + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists()) + require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="function")`).Exists()) + require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists()) + instructions := gjson.GetBytes(upstream.lastBody, "instructions").String() + require.NotContains(t, instructions, "image_generation") +} + func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index 7081653d80..09472fbaf1 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -760,6 +760,8 @@ func (s *OpenAIGatewayService) buildOpenAIImagesRequest( if strings.TrimSpace(contentType) != "" { req.Header.Set("Content-Type", contentType) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) return req, nil } diff --git a/backend/internal/service/openai_model_alias.go b/backend/internal/service/openai_model_alias.go index ac2a8cf942..4e3d3b2d9a 100644 --- a/backend/internal/service/openai_model_alias.go +++ b/backend/internal/service/openai_model_alias.go @@ -65,6 +65,12 @@ func normalizeKnownOpenAICodexModel(model string) string { } switch { + case strings.Contains(normalized, "gpt-5.6-sol"): + return "gpt-5.6-sol" + case strings.Contains(normalized, "gpt-5.6-terra"): + return "gpt-5.6-terra" + case strings.Contains(normalized, "gpt-5.6-luna"): + return "gpt-5.6-luna" case strings.Contains(normalized, "gpt-5.5-pro"): return "gpt-5.5-pro" case strings.Contains(normalized, "gpt-5.5"): diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index 6515c0c4e5..a507213701 100644 --- a/backend/internal/service/openai_tool_continuation.go +++ b/backend/internal/service/openai_tool_continuation.go @@ -215,6 +215,84 @@ func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValid return result } +// ToolCallOutputContextCoverage 描述 input 中工具输出与可重建上下文的覆盖关系, +// 用于判断剥离 previous_response_id 后上游能否仅凭 input 重建工具续链。 +type ToolCallOutputContextCoverage struct { + HasFunctionCallOutput bool + // ContextCoversAllCallIDs 表示每个工具输出的 call_id 都能在 input 内找到 + // 同 call_id 的工具调用上下文项或同 id 的 item_reference,且不存在缺失 call_id 的输出。 + // 任一输出无法由 input 自身重建时为 false,此时剥离 previous_response_id 会导致 + // 上游以 "No tool call found for function call output" 拒绝请求。 + ContextCoversAllCallIDs bool +} + +// AnalyzeToolCallOutputContextCoverageBytes 全量扫描 input,按 call_id 精确匹配工具输出 +// 与可重建上下文。不能复用 ValidateFunctionCallOutputContextBytes 的 HasToolCallContext: +// 该标志只代表"存在某一个上下文项",部分覆盖的续链仍会被上游拒绝。 +func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContextCoverage { + coverage := ToolCallOutputContextCoverage{} + if len(body) == 0 { + return coverage + } + input := parseRawJSONView(body).Get("input") + if !input.IsArray() { + return coverage + } + + missingCallID := false + var outputCallIDs map[string]struct{} + var contextIDs map[string]struct{} + input.ForEach(func(_, item gjson.Result) bool { + if !item.IsObject() { + return true + } + itemType := item.Get("type").String() + switch { + case isCodexToolCallOutputItemType(itemType): + coverage.HasFunctionCallOutput = true + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + missingCallID = true + return true + } + if outputCallIDs == nil { + outputCallIDs = make(map[string]struct{}) + } + outputCallIDs[callID] = struct{}{} + case isCodexToolCallContextItemType(itemType): + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[callID] = struct{}{} + case itemType == "item_reference": + idValue := strings.TrimSpace(item.Get("id").String()) + if idValue == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[idValue] = struct{}{} + } + return true + }) + + if !coverage.HasFunctionCallOutput || missingCallID { + return coverage + } + for callID := range outputCallIDs { + if _, ok := contextIDs[callID]; !ok { + return coverage + } + } + coverage.ContextCoversAllCallIDs = true + return coverage +} + // ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果: // 1) 无工具输出直接返回 // 2) 若已存在工具调用上下文则提前返回 diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go index 4610652b6c..569d89eff0 100644 --- a/backend/internal/service/openai_tool_continuation_test.go +++ b/backend/internal/service/openai_tool_continuation_test.go @@ -184,3 +184,109 @@ func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T }) } } + +func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) { + cases := []struct { + name string + body map[string]any + hasOutput bool + coversAllIDs bool + }{ + { + name: "no_input", + body: map[string]any{"model": "gpt-5.1"}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "no_tool_output", + body: map[string]any{"input": []any{ + map[string]any{"type": "message", "content": "hi"}, + }}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "all_outputs_covered_by_context", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_outputs_covered_by_item_reference", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "item_reference", "id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + // 关键回归用例:input 内存在某一个上下文项,但另一个输出的 call_id + // 只能由上游会话链(previous_response_id)解析——不可剥离。 + name: "partial_coverage_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "unrelated_context_does_not_cover", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_x"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "output_missing_call_id_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "mixed_context_and_reference_cover_all", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + map[string]any{"type": "item_reference", "id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_codex_output_types_covered", + body: map[string]any{"input": []any{ + map[string]any{"type": "tool_search_output", "call_id": "call_s"}, + map[string]any{"type": "tool_search_call", "call_id": "call_s"}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "call_m"}, + map[string]any{"type": "mcp_tool_call", "call_id": "call_m"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + } + + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + bodyBytes, err := json.Marshal(tt.body) + require.NoError(t, err) + + coverage := AnalyzeToolCallOutputContextCoverageBytes(bodyBytes) + require.Equal(t, tt.hasOutput, coverage.HasFunctionCallOutput, "HasFunctionCallOutput") + require.Equal(t, tt.coversAllIDs, coverage.ContextCoversAllCallIDs, "ContextCoversAllCallIDs") + }) + } +} diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 067eeb6029..bbca9776ab 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -1183,6 +1183,10 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( headers.Set("user-agent", codexCLIUserAgent) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。 + // 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。 + account.ApplyHeaderOverrides(headers) + return headers, sessionResolution, nil } @@ -2448,11 +2452,15 @@ func stripCodexSparkImageGenerationToolFromRawPayload(payload []byte, model stri if !isCodexSparkModel(model) || !openAIRequestBodyHasImageGenerationTool(payload) { return payload, false, nil } + return stripOpenAIImageGenerationToolFromRawPayload(payload) +} + +func stripOpenAIImageGenerationToolFromRawPayload(payload []byte) ([]byte, bool, error) { payloadMap := make(map[string]any) if err := json.Unmarshal(payload, &payloadMap); err != nil { return payload, false, err } - if !stripCodexSparkImageGenerationTools(payloadMap) { + if !stripOpenAIImageGenerationTools(payloadMap) { return payload, false, nil } rebuilt, err := json.Marshal(payloadMap) @@ -2671,7 +2679,11 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } apiKey := getAPIKeyFromContext(c) imageGenerationAllowed := GroupAllowsImageGeneration(apiKeyGroup(apiKey)) - codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) + codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow + if isCodexCLI { + codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() + } + codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) if codexBridgeEnabled { payloadMap := make(map[string]any) if err := json.Unmarshal(normalized, &payloadMap); err != nil { @@ -2709,6 +2721,14 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } normalized = next } + if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { + if stripped, changed, stripErr := stripOpenAIImageGenerationToolFromRawPayload(normalized); stripErr != nil { + return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr) + } else if changed { + normalized = stripped + logOpenAIWSModeInfo("ingress_ws_codex_image_tool_stripped_by_policy account_id=%d", account.ID) + } + } if stripped, changed, stripErr := stripCodexSparkImageGenerationToolFromRawPayload(normalized, upstreamModel); stripErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr) } else if changed { @@ -4309,87 +4329,8 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability( if s == nil { return nil, nil } - responseID := strings.TrimSpace(previousResponseID) - if responseID == "" { - return nil, nil - } - store := s.getOpenAIWSStateStore() - if store == nil { - return nil, nil - } - - accountID, err := store.GetResponseAccount(ctx, derefGroupID(groupID), responseID) - if err != nil || accountID <= 0 { - return nil, nil - } - if excludedIDs != nil { - if _, excluded := excludedIDs[accountID]; excluded { - return nil, nil - } - } - - account, err := s.getSchedulableAccount(ctx, accountID) - if err != nil || account == nil { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - // 非 WSv2 场景(如 force_http/全局关闭)不应使用 previous_response_id 粘连, - // 以保持“回滚到 HTTP”后的历史行为一致性。 - if s.getOpenAIWSProtocolResolver().Resolve(account).Transport != OpenAIUpstreamTransportResponsesWebsocketV2 { - return nil, nil - } - if shouldClearStickySession(account, requestedModel) || !account.IsOpenAI() || !account.IsSchedulable() { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - if requestedModel != "" && !account.IsModelSupported(requestedModel) { - return nil, nil - } - if !account.SupportsOpenAIEndpointCapability(requiredCapability) { - return nil, nil - } - // Quota auto-pause must also gate the previous_response_id sticky path; otherwise an - // account over its 5h/7d threshold keeps serving the same response chain even though - // normal scheduling skips it. Pause is transient, so fall through to normal scheduling - // without deleting the binding (the window may reset before the next turn). - if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { - return nil, nil - } - if s.schedulerSnapshot != nil && s.accountRepo != nil { - latest, latestErr := s.accountRepo.GetByID(ctx, account.ID) - if latestErr != nil || latest == nil { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - if requestedModel != "" && !latest.IsModelSupported(requestedModel) { - return nil, nil - } - if !latest.SupportsOpenAIEndpointCapability(requiredCapability) { - return nil, nil - } - if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused { - return nil, nil - } - if s.isOpenAIAccountRuntimeBlocked(latest) { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) - return nil, nil - } - account = latest - } - if requireCompact && openAICompactSupportTier(account) == 0 { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + accountID, account, responseID, store := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact) + if accountID <= 0 || account == nil || store == nil { return nil, nil } @@ -4423,6 +4364,117 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability( return nil, nil } +func (s *OpenAIGatewayService) ResolveAccountIDByPreviousResponseIDForScheduler( + ctx context.Context, + groupID *int64, + previousResponseID string, + requestedModel string, + excludedIDs map[int64]struct{}, + requiredCapability OpenAIEndpointCapability, + requireCompact bool, +) int64 { + accountID, _, _, _ := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact) + return accountID +} + +func (s *OpenAIGatewayService) resolveAccountByPreviousResponseIDForCapability( + ctx context.Context, + groupID *int64, + previousResponseID string, + requestedModel string, + excludedIDs map[int64]struct{}, + requiredCapability OpenAIEndpointCapability, + requireCompact bool, +) (int64, *Account, string, OpenAIWSStateStore) { + if s == nil { + return 0, nil, "", nil + } + responseID := strings.TrimSpace(previousResponseID) + if responseID == "" { + return 0, nil, "", nil + } + store := s.getOpenAIWSStateStore() + if store == nil { + return 0, nil, "", nil + } + + accountID, err := store.GetResponseAccount(ctx, derefGroupID(groupID), responseID) + if err != nil || accountID <= 0 { + return 0, nil, "", nil + } + if excludedIDs != nil { + if _, excluded := excludedIDs[accountID]; excluded { + return 0, nil, "", nil + } + } + + account, err := s.getSchedulableAccount(ctx, accountID) + if err != nil || account == nil { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + // 非 WSv2 场景(如 force_http/全局关闭)不应使用 previous_response_id 粘连, + // 以保持“回滚到 HTTP”后的历史行为一致性。 + if s.getOpenAIWSProtocolResolver().Resolve(account).Transport != OpenAIUpstreamTransportResponsesWebsocketV2 { + return 0, nil, "", nil + } + if shouldClearStickySession(account, requestedModel) || !account.IsOpenAI() || !account.IsSchedulable() { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + if requestedModel != "" && !account.IsModelSupported(requestedModel) { + return 0, nil, "", nil + } + if !account.SupportsOpenAIEndpointCapability(requiredCapability) { + return 0, nil, "", nil + } + // Quota auto-pause must also gate the previous_response_id sticky path; otherwise an + // account over its 5h/7d threshold keeps serving the same response chain even though + // normal scheduling skips it. Pause is transient, so fall through to normal scheduling + // without deleting the binding (the window may reset before the next turn). + if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { + return 0, nil, "", nil + } + if s.schedulerSnapshot != nil && s.accountRepo != nil { + latest, latestErr := s.accountRepo.GetByID(ctx, account.ID) + if latestErr != nil || latest == nil { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + if requestedModel != "" && !latest.IsModelSupported(requestedModel) { + return 0, nil, "", nil + } + if !latest.SupportsOpenAIEndpointCapability(requiredCapability) { + return 0, nil, "", nil + } + if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused { + return 0, nil, "", nil + } + if s.isOpenAIAccountRuntimeBlocked(latest) { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + account = latest + } + if requireCompact && openAICompactSupportTier(account) == 0 { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return 0, nil, "", nil + } + return accountID, account, responseID, store +} + func classifyOpenAIWSAcquireError(err error) string { if err == nil { return "acquire_conn" diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index 0d19a189b0..ca7c36aaa7 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -169,6 +169,26 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) { }) } +func TestStripOpenAIImageGenerationToolFromRawPayload(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.4", + "tools":[ + {"type":"function","name":"shell"}, + {"type":"image_generation","output_format":"png"} + ], + "tool_choice":{"type":"image_generation"} + }`) + + updated, changed, err := stripOpenAIImageGenerationToolFromRawPayload(payload) + + require.NoError(t, err) + require.True(t, changed) + require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists()) + require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) +} + func TestAlignStoreDisabledPreviousResponseID(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go index 4fc6a9266e..e33dcf82a8 100644 --- a/backend/internal/service/ops_models.go +++ b/backend/internal/service/ops_models.go @@ -1,6 +1,9 @@ package service -import "time" +import ( + "strings" + "time" +) type OpsSystemLog struct { ID int64 `json:"id"` @@ -65,17 +68,22 @@ type OpsErrorLog struct { RequestedModel string `json:"requested_model"` UpstreamModel string `json:"upstream_model"` RequestType *int16 `json:"request_type"` + UserAgent string `json:"user_agent"` // 关联 api_key 名称(LEFT JOIN api_keys 取得;软删只覆盖 key 列,name 保留,故已删 key 仍有原名)。 APIKeyName string `json:"api_key_name,omitempty"` APIKeyDeleted bool `json:"api_key_deleted,omitempty"` + + // 已删除 KEY 所有者(INVALID_API_KEY 且该 key 曾存在时的归因快照)。 + // 认证失败行 user_id 为空,列表用户列以此回退显示所有者。 + DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"` + DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"` } type OpsErrorLogDetail struct { OpsErrorLog ErrorBody string `json:"error_body"` - UserAgent string `json:"user_agent"` // Upstream context (optional) UpstreamStatusCode *int `json:"upstream_status_code,omitempty"` @@ -93,11 +101,10 @@ type OpsErrorLogDetail struct { // vNext metric semantics IsBusinessLimited bool `json:"is_business_limited"` - // Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted) - AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"` - DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"` - DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"` - DeletedKeyName string `json:"deleted_key_name,omitempty"` + // Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted). + // OwnerUserID/OwnerEmail 已上移到 OpsErrorLog(列表用户列回退需要)。 + AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"` + DeletedKeyName string `json:"deleted_key_name,omitempty"` // Bound (non-deleted) key prefix, snapshotted at error time; mutually exclusive with AttemptedKeyPrefix. APIKeyPrefix string `json:"api_key_prefix,omitempty"` @@ -142,8 +149,14 @@ type OpsErrorLogFilter struct { // ExcludeCountTokens drops count_tokens probe errors (is_count_tokens=true). ExcludeCountTokens bool + // IncludeRecoveredUpstream 显式豁免 status>=400 守卫(仅在 Phase=="upstream" 时生效): + // ops 专用上游错误列表需要看到 status<400 的 recovered upstream 行。 + // 请求错误语义的端点不设此开关,phase=upstream 过滤照常生效且守卫保留。 + IncludeRecoveredUpstream bool + // ErrorPhasesAny / ErrorTypesAny add plain ANY() filters WITHOUT touching the - // special-cased single `Phase` field (only Phase=="upstream" bypasses the status>=400 clause). + // special-cased single `Phase` field (only Phase=="upstream" with + // IncludeRecoveredUpstream bypasses the status>=400 clause). // NOTE: these ANY filters do NOT bypass status>=400; records with error_phase='upstream' // but status_code<400 (recovered upstream errors) remain excluded. // Used to map user-facing coarse categories to backend conditions. @@ -158,6 +171,19 @@ type OpsErrorLogFilter struct { Page int PageSize int + + // SortBy/SortOrder: server-side sorting aligned with the usage-log list. + // Repo whitelists columns (created_at/model/status_code); anything else + // falls back to created_at. SortOrder is "asc"/"desc" (default desc). + SortBy string + SortOrder string +} + +// SetSort normalizes raw sort_by/sort_order query values into the filter. +// Shared by the admin and user-facing error list handlers. +func (f *OpsErrorLogFilter) SetSort(sortBy, sortOrder string) { + f.SortBy = strings.TrimSpace(sortBy) + f.SortOrder = strings.TrimSpace(sortOrder) } type OpsErrorLogList struct { diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go index a8c8a4bb5c..61f85ef904 100644 --- a/backend/internal/service/ops_service.go +++ b/backend/internal/service/ops_service.go @@ -359,10 +359,12 @@ func (s *OpsService) ListUserErrorRequests(ctx context.Context, userID int64, fi filter.UserQuery = "" filter.Owner = "" filter.Source = "" - // 清空 Phase 是防御:Phase 是单值特殊字段,仅当其 == "upstream" 时 buildOpsErrorLogsWhere 才跳过 status>=400 子句。 - // 用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny(纯 ANY 过滤,不影响 status>=400 子句), - // 因此 recovered upstream(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。 + // 清空 Phase 是防御:用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny + //(纯 ANY 过滤,不影响 status>=400 子句)。守卫豁免现在还需要 + // IncludeRecoveredUpstream(用户端永不设置),recovered upstream + //(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。 filter.Phase = "" + filter.IncludeRecoveredUpstream = false list, err := s.opsRepo.ListErrorLogs(ctx, filter) if err != nil { diff --git a/backend/internal/service/ops_service_user_error_test.go b/backend/internal/service/ops_service_user_error_test.go index 9027ff0788..c3b0967b67 100644 --- a/backend/internal/service/ops_service_user_error_test.go +++ b/backend/internal/service/ops_service_user_error_test.go @@ -184,16 +184,16 @@ func TestGetUserErrorRequestDetail_DeletedKeyOwnerAccess(t *testing.T) { mk := func() *OpsErrorLogDetail { return &OpsErrorLogDetail{ OpsErrorLog: OpsErrorLog{ - ID: 55, - Phase: "auth", - Type: "api_error", - StatusCode: 401, - Message: "Invalid API key", - UserID: nil, - APIKeyName: "my-old-key", - APIKeyDeleted: true, + ID: 55, + Phase: "auth", + Type: "api_error", + StatusCode: 401, + Message: "Invalid API key", + UserID: nil, + APIKeyName: "my-old-key", + APIKeyDeleted: true, + DeletedKeyOwnerUserID: &ownerUID, }, - DeletedKeyOwnerUserID: &ownerUID, } } diff --git a/backend/internal/service/ops_user_error.go b/backend/internal/service/ops_user_error.go index 7dd128afa7..e3055c2392 100644 --- a/backend/internal/service/ops_user_error.go +++ b/backend/internal/service/ops_user_error.go @@ -3,9 +3,12 @@ package service import "time" // UserErrorRequest 是面向终端用户的"错误请求"精简脱敏视图(白名单)。 -// 严禁包含 client_ip / user_agent / account / api_key_prefix / upstream_endpoint / -// user_email 等敏感或内部字段。注:message(网关标准化错误描述)与 key_name +// 严禁包含 account / api_key_prefix / upstream_endpoint / user_email 等 +// 敏感或内部字段。注:message(网关标准化错误描述)与 key_name // (用户自有 API Key 名称,KeysView 中本就可见)经产品决策对该用户开放; +// client_ip / user_agent / group_name / request_type / stream 均为该用户 +// 自己请求的属性,经产品决策(2026-07-03)开放, +// 与用量明细已向用户展示自身 ip_address/user_agent/分组/类型 的口径对齐; // error_body 仅在详情接口(GetUserErrorRequestDetail)按归属校验后返回。 type UserErrorRequest struct { ID int64 `json:"id"` @@ -18,6 +21,11 @@ type UserErrorRequest struct { Message string `json:"message"` KeyName string `json:"key_name"` KeyDeleted bool `json:"key_deleted"` + ClientIP string `json:"client_ip,omitempty"` + GroupName string `json:"group_name,omitempty"` + RequestType *int16 `json:"request_type,omitempty"` + Stream bool `json:"stream"` + UserAgent string `json:"user_agent,omitempty"` } // UserErrorRequestList 是用户错误请求分页结果。 @@ -90,6 +98,10 @@ func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest { if model == "" { model = e.Model } + clientIP := "" + if e.ClientIP != nil { + clientIP = *e.ClientIP + } return &UserErrorRequest{ ID: e.ID, CreatedAt: e.CreatedAt, @@ -101,6 +113,11 @@ func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest { Message: e.Message, KeyName: e.APIKeyName, KeyDeleted: e.APIKeyDeleted, + ClientIP: clientIP, + GroupName: e.GroupName, + RequestType: e.RequestType, + Stream: e.Stream, + UserAgent: e.UserAgent, } } diff --git a/backend/internal/service/ops_user_error_test.go b/backend/internal/service/ops_user_error_test.go index 31b0c26933..9e0bc164b4 100644 --- a/backend/internal/service/ops_user_error_test.go +++ b/backend/internal/service/ops_user_error_test.go @@ -122,9 +122,11 @@ func TestToUserErrorRequestDetail_WhitelistAndRedacts(t *testing.T) { UserEmail: "secret@example.com", ClientIP: func() *string { s := "1.2.3.4"; return &s }(), UpstreamEndpoint: "https://api.openai.com/v1/chat/completions", + UserAgent: "codex_cli_rs/0.125.0", + GroupName: "grp-a", + Stream: true, }, ErrorBody: `{"error":{"message":"upstream failed","type":"server_error"}}`, - UserAgent: "Mozilla/5.0 secret-agent", UpstreamStatusCode: &upstreamStatus, } @@ -147,13 +149,27 @@ func TestToUserErrorRequestDetail_WhitelistAndRedacts(t *testing.T) { t.Errorf("UpstreamStatusCode mismatch") } + // client_ip / user_agent / group_name / stream 经产品决策开放(与用量明细口径对齐) + if out.ClientIP != "1.2.3.4" { + t.Errorf("want client_ip=1.2.3.4, got %q", out.ClientIP) + } + if out.UserAgent != "codex_cli_rs/0.125.0" { + t.Errorf("want user_agent=codex_cli_rs/0.125.0, got %q", out.UserAgent) + } + if out.GroupName != "grp-a" { + t.Errorf("want group_name=grp-a, got %q", out.GroupName) + } + if !out.Stream { + t.Errorf("want stream=true") + } + // 序列化后不含敏感字段 b, err := json.Marshal(out) if err != nil { t.Fatalf("json.Marshal failed: %v", err) } raw := string(b) - for _, forbidden := range []string{"user_email", "client_ip", "upstream_endpoint", "user_agent"} { + for _, forbidden := range []string{"user_email", "upstream_endpoint"} { if strings.Contains(raw, forbidden) { t.Errorf("sensitive field %q leaked in JSON output: %s", forbidden, raw) } diff --git a/backend/internal/service/payment_amounts.go b/backend/internal/service/payment_amounts.go index a7f620d33e..2fd00c5957 100644 --- a/backend/internal/service/payment_amounts.go +++ b/backend/internal/service/payment_amounts.go @@ -16,6 +16,15 @@ func normalizeBalanceRechargeMultiplier(multiplier float64) float64 { return multiplier } +// normalizeSubscriptionUSDToCNYRate 将非法值归一为 0(换算关闭)。 +// 与余额倍率不同,0 是合法状态:表示订阅保持 price 直付的存量行为。 +func normalizeSubscriptionUSDToCNYRate(rate float64) float64 { + if math.IsNaN(rate) || math.IsInf(rate, 0) || rate < 0 { + return 0 + } + return rate +} + func calculateCreditedBalance(paymentAmount, multiplier float64) float64 { return decimal.NewFromFloat(paymentAmount). Mul(decimal.NewFromFloat(normalizeBalanceRechargeMultiplier(multiplier))). diff --git a/backend/internal/service/payment_config_limits.go b/backend/internal/service/payment_config_limits.go index 45b24bfce7..202eea9f26 100644 --- a/backend/internal/service/payment_config_limits.go +++ b/backend/internal/service/payment_config_limits.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "strings" dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/paymentproviderinstance" @@ -31,6 +32,7 @@ func (s *PaymentConfigService) GetAvailableMethodLimits(ctx context.Context) (*M continue } ml := pcAggregateMethodLimits(pt, insts) + ml.DisplayName = s.pcAggregateMethodDisplayName(pt, insts) ml.Currency = currency resp.Methods[ml.PaymentType] = ml } @@ -93,6 +95,7 @@ func (s *PaymentConfigService) GetMethodLimits(ctx context.Context, types []stri continue } ml := pcAggregateMethodLimits(pt, matching) + ml.DisplayName = s.pcAggregateMethodDisplayName(pt, matching) ml.Currency = currency result = append(result, ml) } @@ -163,6 +166,53 @@ func (s *PaymentConfigService) pcInstancePaymentCurrency(inst *dbent.PaymentProv return paymentProviderConfigCurrency(inst.ProviderKey, cfg) } +type easyPayCustomMethodDisplayConfig struct { + Type string `json:"type"` + DisplayName string `json:"displayName"` +} + +func (s *PaymentConfigService) pcAggregateMethodDisplayName(pt string, instances []*dbent.PaymentProviderInstance) string { + pt = strings.TrimSpace(pt) + if pt == "" { + return "" + } + for _, inst := range instances { + displayName := s.pcInstanceEasyPayCustomMethodDisplayName(inst, pt) + if displayName != "" { + return displayName + } + } + return "" +} + +func (s *PaymentConfigService) pcInstanceEasyPayCustomMethodDisplayName(inst *dbent.PaymentProviderInstance, pt string) string { + if inst == nil || inst.ProviderKey != payment.TypeEasyPay { + return "" + } + cfg := map[string]string{} + if s != nil { + decrypted, err := s.decryptConfig(inst.Config) + if err == nil && decrypted != nil { + cfg = decrypted + } + } + raw := strings.TrimSpace(cfg["customMethods"]) + if raw == "" { + return "" + } + + var methods []easyPayCustomMethodDisplayConfig + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return "" + } + for _, method := range methods { + if strings.TrimSpace(method.Type) == pt { + return strings.TrimSpace(method.DisplayName) + } + } + return "" +} + // pcGroupByPaymentType groups instances by user-facing payment type. // For Stripe providers, ALL sub-types (card, link, alipay, wxpay) map to "stripe" // because the user sees a single "Stripe" button, not individual sub-methods. diff --git a/backend/internal/service/payment_config_limits_test.go b/backend/internal/service/payment_config_limits_test.go index c0aa2b27a5..a70bc90a29 100644 --- a/backend/internal/service/payment_config_limits_test.go +++ b/backend/internal/service/payment_config_limits_test.go @@ -255,6 +255,28 @@ func TestGetAvailableMethodLimitsOmitsMixedCurrencyMethod(t *testing.T) { require.Equal(t, "PAYMENT_METHOD_CURRENCY_CONFLICT", appErr.Reason) } +func TestGetAvailableMethodLimitsIncludesEasyPayCustomMethodDisplayName(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + + _, err := client.PaymentProviderInstance.Create(). + SetProviderKey(payment.TypeEasyPay). + SetName("EasyPay Custom"). + SetConfig(`{"customMethods":"[{\"type\":\"ldc\",\"upstreamType\":\"ldc\",\"displayName\":\"LDC Pay\"}]"}`). + SetSupportedTypes("alipay,wxpay,ldc"). + SetEnabled(true). + Save(ctx) + require.NoError(t, err) + + svc := &PaymentConfigService{entClient: client} + resp, err := svc.GetAvailableMethodLimits(ctx) + require.NoError(t, err) + + limits, ok := resp.Methods["ldc"] + require.True(t, ok, "expected custom EasyPay method limits to be visible") + require.Equal(t, "LDC Pay", limits.DisplayName) +} + func TestPcComputeGlobalRange(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_config_providers.go b/backend/internal/service/payment_config_providers.go index 7e92558568..d1bf2de7aa 100644 --- a/backend/internal/service/payment_config_providers.go +++ b/backend/internal/service/payment_config_providers.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "log/slog" + "regexp" "strconv" "strings" @@ -185,6 +186,11 @@ func (s *PaymentConfigService) CreateProviderInstance(ctx context.Context, req C if err := validateProviderRequest(req.ProviderKey, req.Name, typesStr); err != nil { return nil, err } + if req.ProviderKey == payment.TypeEasyPay { + if err := validateEasyPayCustomMethods(req.Config, typesStr); err != nil { + return nil, err + } + } if err := s.validateVisibleMethodEnablementConflicts(ctx, 0, req.ProviderKey, typesStr, req.Enabled); err != nil { return nil, err } @@ -217,6 +223,67 @@ func validateProviderRequest(providerKey, name, supportedTypes string) error { return nil } +var easyPayCustomMethodCodePattern = regexp.MustCompile(`^[a-z0-9_-]+$`) + +type easyPayCustomMethodConfig struct { + Type string `json:"type"` + UpstreamType string `json:"upstreamType"` + DisplayName string `json:"displayName"` +} + +func validateEasyPayCustomMethods(config map[string]string, supportedTypes string) error { + if config == nil { + config = map[string]string{} + } + raw := strings.TrimSpace(config["customMethods"]) + methods := make([]easyPayCustomMethodConfig, 0) + if raw != "" { + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods must be a JSON array") + } + } + + customTypes := make(map[string]struct{}, len(methods)) + for _, method := range methods { + method.Type = strings.TrimSpace(method.Type) + method.UpstreamType = strings.TrimSpace(method.UpstreamType) + if method.Type == "" || method.UpstreamType == "" { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods upstreamType is required") + } + if !easyPayCustomMethodCodePattern.MatchString(method.Type) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods type may only contain lowercase letters, digits, underscores, and hyphens") + } + if !easyPayCustomMethodCodePattern.MatchString(method.UpstreamType) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods upstreamType may only contain lowercase letters, digits, underscores, and hyphens") + } + if easyPayCustomMethodTypeConflictsWithBuiltin(method.Type) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods type cannot start with alipay or wxpay") + } + if _, exists := customTypes[method.Type]; exists { + return infraerrors.BadRequest("VALIDATION_ERROR", "duplicate customMethods type") + } + customTypes[method.Type] = struct{}{} + } + + for _, supportedType := range splitTypes(supportedTypes) { + supportedType = strings.TrimSpace(supportedType) + if supportedType == "" || supportedType == payment.TypeAlipay || supportedType == payment.TypeWxpay { + continue + } + if !easyPayCustomMethodCodePattern.MatchString(supportedType) { + return infraerrors.BadRequest("VALIDATION_ERROR", fmt.Sprintf("supported EasyPay custom type %s may only contain lowercase letters, digits, underscores, and hyphens", supportedType)) + } + if _, exists := customTypes[supportedType]; !exists { + return infraerrors.BadRequest("VALIDATION_ERROR", fmt.Sprintf("supported EasyPay custom type %s has no customMethods mapping", supportedType)) + } + } + return nil +} + +func easyPayCustomMethodTypeConflictsWithBuiltin(methodType string) bool { + return strings.HasPrefix(methodType, payment.TypeAlipay) || strings.HasPrefix(methodType, payment.TypeWxpay) +} + // UpdateProviderInstance updates a provider instance by ID (patch semantics). // NOTE: This function exceeds 30 lines due to per-field nil-check patch update // boilerplate and pending-order safety checks. @@ -279,6 +346,18 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in WithMetadata(map[string]string{"count": strconv.Itoa(count)}) } } + configToValidate := mergedConfig + if configToValidate == nil { + configToValidate, err = s.decryptConfig(current.Config) + if err != nil { + return nil, fmt.Errorf("decrypt existing config: %w", err) + } + } + if current.ProviderKey == payment.TypeEasyPay { + if err := validateEasyPayCustomMethods(configToValidate, nextSupportedTypes); err != nil { + return nil, err + } + } // Validate merged config when the instance will end up enabled. // This surfaces provider-level errors (e.g. wxpay missing certSerial) at save time, // so admins see them in the dialog instead of only when an order is created. @@ -287,13 +366,6 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in finalEnabled = *req.Enabled } if finalEnabled { - configToValidate := mergedConfig - if configToValidate == nil { - configToValidate, err = s.decryptConfig(current.Config) - if err != nil { - return nil, fmt.Errorf("decrypt existing config: %w", err) - } - } if err := s.validateProviderConfig(current.ProviderKey, configToValidate); err != nil { return nil, err } diff --git a/backend/internal/service/payment_config_providers_test.go b/backend/internal/service/payment_config_providers_test.go index 43708de73d..74fd2a3467 100644 --- a/backend/internal/service/payment_config_providers_test.go +++ b/backend/internal/service/payment_config_providers_test.go @@ -114,6 +114,92 @@ func TestValidateProviderRequest(t *testing.T) { } } +func TestValidateEasyPayCustomMethods(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + config map[string]string + supportedTypes string + wantErr string + }{ + { + name: "valid custom methods", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"}]`}, + supportedTypes: "alipay,wxpay,ldc", + }, + { + name: "malformed custom methods json", + config: map[string]string{"customMethods": `not-json`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods must be a JSON array", + }, + { + name: "missing upstream type", + config: map[string]string{"customMethods": `[{"type":"ldc","displayName":"LDC"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods upstreamType is required", + }, + { + name: "duplicate custom type", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"},{"type":"ldc","upstreamType":"epay2"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "duplicate customMethods type", + }, + { + name: "custom type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"LDC","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods type may only contain lowercase letters", + }, + { + name: "upstream type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"ALIPAY"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods upstreamType may only contain lowercase letters", + }, + { + name: "custom type uses alipay prefix", + config: map[string]string{"customMethods": `[{"type":"alipay_hk","upstreamType":"hkpay"}]`}, + supportedTypes: "alipay,wxpay,alipay_hk", + wantErr: "customMethods type cannot start with alipay or wxpay", + }, + { + name: "custom type uses wxpay prefix", + config: map[string]string{"customMethods": `[{"type":"wxpay_usdt","upstreamType":"usdt"}]`}, + supportedTypes: "alipay,wxpay,wxpay_usdt", + wantErr: "customMethods type cannot start with alipay or wxpay", + }, + { + name: "supported custom type missing mapping", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,ldc,usdt_trc20", + wantErr: "supported EasyPay custom type usdt_trc20 has no customMethods mapping", + }, + { + name: "supported custom type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,LDC", + wantErr: "supported EasyPay custom type LDC may only contain lowercase letters", + }, + } + + for _, tc := range tests { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + err := validateEasyPayCustomMethods(tc.config, tc.supportedTypes) + if tc.wantErr == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + require.Contains(t, err.Error(), tc.wantErr) + }) + } +} + func TestIsSensitiveProviderConfigField(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_config_service.go b/backend/internal/service/payment_config_service.go index 022b1b0156..52a7ddc67e 100644 --- a/backend/internal/service/payment_config_service.go +++ b/backend/internal/service/payment_config_service.go @@ -24,17 +24,20 @@ const ( SettingLoadBalanceStrategy = "LOAD_BALANCE_STRATEGY" SettingBalancePayDisabled = "BALANCE_PAYMENT_DISABLED" SettingBalanceRechargeMult = "BALANCE_RECHARGE_MULTIPLIER" - SettingRechargeFeeRate = "RECHARGE_FEE_RATE" - SettingProductNamePrefix = "PRODUCT_NAME_PREFIX" - SettingProductNameSuffix = "PRODUCT_NAME_SUFFIX" - SettingHelpImageURL = "PAYMENT_HELP_IMAGE_URL" - SettingHelpText = "PAYMENT_HELP_TEXT" - SettingCancelRateLimitOn = "CANCEL_RATE_LIMIT_ENABLED" - SettingCancelRateLimitMax = "CANCEL_RATE_LIMIT_MAX" - SettingCancelWindowSize = "CANCEL_RATE_LIMIT_WINDOW" - SettingCancelWindowUnit = "CANCEL_RATE_LIMIT_UNIT" - SettingCancelWindowMode = "CANCEL_RATE_LIMIT_WINDOW_MODE" - SettingAlipayForceQRCode = "ALIPAY_FORCE_QRCODE" + // SettingSubscriptionUSDToCNYRate 是订阅 CNY 换算汇率(1 USD = X CNY)。 + // 0/未配置 = 关闭换算(订阅按 price 数值直付),显式配置后 CNY 通道订阅按 price × rate 收款。 + SettingSubscriptionUSDToCNYRate = "SUBSCRIPTION_USD_TO_CNY_RATE" + SettingRechargeFeeRate = "RECHARGE_FEE_RATE" + SettingProductNamePrefix = "PRODUCT_NAME_PREFIX" + SettingProductNameSuffix = "PRODUCT_NAME_SUFFIX" + SettingHelpImageURL = "PAYMENT_HELP_IMAGE_URL" + SettingHelpText = "PAYMENT_HELP_TEXT" + SettingCancelRateLimitOn = "CANCEL_RATE_LIMIT_ENABLED" + SettingCancelRateLimitMax = "CANCEL_RATE_LIMIT_MAX" + SettingCancelWindowSize = "CANCEL_RATE_LIMIT_WINDOW" + SettingCancelWindowUnit = "CANCEL_RATE_LIMIT_UNIT" + SettingCancelWindowMode = "CANCEL_RATE_LIMIT_WINDOW_MODE" + SettingAlipayForceQRCode = "ALIPAY_FORCE_QRCODE" ) // Default values for payment configuration settings. @@ -54,13 +57,15 @@ type PaymentConfig struct { EnabledTypes []string `json:"enabled_payment_types"` BalanceDisabled bool `json:"balance_disabled"` BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` - RechargeFeeRate float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy string `json:"load_balance_strategy"` - ProductNamePrefix string `json:"product_name_prefix"` - ProductNameSuffix string `json:"product_name_suffix"` - HelpImageURL string `json:"help_image_url"` - HelpText string `json:"help_text"` - StripePublishableKey string `json:"stripe_publishable_key,omitempty"` + // SubscriptionUSDToCNYRate 为 0 时订阅换算关闭(兼容存量行为)。 + SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` + RechargeFeeRate float64 `json:"recharge_fee_rate"` + LoadBalanceStrategy string `json:"load_balance_strategy"` + ProductNamePrefix string `json:"product_name_prefix"` + ProductNameSuffix string `json:"product_name_suffix"` + HelpImageURL string `json:"help_image_url"` + HelpText string `json:"help_text"` + StripePublishableKey string `json:"stripe_publishable_key,omitempty"` // Cancel rate limit settings CancelRateLimitEnabled bool `json:"cancel_rate_limit_enabled"` @@ -84,6 +89,7 @@ type UpdatePaymentConfigRequest struct { EnabledTypes []string `json:"enabled_payment_types"` BalanceDisabled *bool `json:"balance_disabled"` BalanceRechargeMultiplier *float64 `json:"balance_recharge_multiplier"` + SubscriptionUSDToCNYRate *float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate *float64 `json:"recharge_fee_rate"` LoadBalanceStrategy *string `json:"load_balance_strategy"` ProductNamePrefix *string `json:"product_name_prefix"` @@ -110,6 +116,7 @@ type UpdatePaymentConfigRequest struct { // MethodLimits holds per-payment-type limits. type MethodLimits struct { PaymentType string `json:"payment_type"` + DisplayName string `json:"display_name,omitempty"` Currency string `json:"currency"` FeeRate float64 `json:"fee_rate"` DailyLimit float64 `json:"daily_limit"` @@ -204,7 +211,7 @@ func (s *PaymentConfigService) GetPaymentConfig(ctx context.Context) (*PaymentCo keys := []string{ SettingPaymentEnabled, SettingMinRechargeAmount, SettingMaxRechargeAmount, SettingDailyRechargeLimit, SettingOrderTimeoutMinutes, SettingMaxPendingOrders, - SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingRechargeFeeRate, SettingLoadBalanceStrategy, + SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingSubscriptionUSDToCNYRate, SettingRechargeFeeRate, SettingLoadBalanceStrategy, SettingProductNamePrefix, SettingProductNameSuffix, SettingHelpImageURL, SettingHelpText, SettingCancelRateLimitOn, SettingCancelRateLimitMax, @@ -233,6 +240,7 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme MaxPendingOrders: pcParseInt(vals[SettingMaxPendingOrders], defaultMaxPendingOrders), BalanceDisabled: vals[SettingBalancePayDisabled] == "true", BalanceRechargeMultiplier: normalizeBalanceRechargeMultiplier(pcParseFloat(vals[SettingBalanceRechargeMult], defaultBalanceRechargeMultiplier)), + SubscriptionUSDToCNYRate: normalizeSubscriptionUSDToCNYRate(pcParseFloat(vals[SettingSubscriptionUSDToCNYRate], 0)), RechargeFeeRate: pcParseFloat(vals[SettingRechargeFeeRate], 0), LoadBalanceStrategy: vals[SettingLoadBalanceStrategy], ProductNamePrefix: vals[SettingProductNamePrefix], @@ -294,6 +302,12 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda return infraerrors.BadRequest("INVALID_BALANCE_RECHARGE_MULTIPLIER", "balance recharge multiplier must be greater than 0") } } + if req.SubscriptionUSDToCNYRate != nil { + v := *req.SubscriptionUSDToCNYRate + if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 { + return infraerrors.BadRequest("INVALID_SUBSCRIPTION_USD_TO_CNY_RATE", "subscription USD to CNY rate must be 0 (disabled) or a positive number") + } + } if req.RechargeFeeRate != nil { v := *req.RechargeFeeRate if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 || v > 100 { @@ -313,6 +327,7 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda SettingMaxPendingOrders: formatPositiveInt(req.MaxPendingOrders), SettingBalancePayDisabled: formatBoolOrEmpty(req.BalanceDisabled), SettingBalanceRechargeMult: formatPositiveFloat(req.BalanceRechargeMultiplier), + SettingSubscriptionUSDToCNYRate: formatPositiveFloatExact(req.SubscriptionUSDToCNYRate), SettingRechargeFeeRate: formatNonNegativeFloat(req.RechargeFeeRate), SettingLoadBalanceStrategy: derefStr(req.LoadBalanceStrategy), SettingProductNamePrefix: derefStr(req.ProductNamePrefix), @@ -352,6 +367,14 @@ func formatPositiveFloat(v *float64) string { return strconv.FormatFloat(*v, 'f', 2, 64) } +// formatPositiveFloatExact 保留完整精度,用于汇率等对小数位敏感的配置。 +func formatPositiveFloatExact(v *float64) string { + if v == nil || *v <= 0 { + return "" // empty → parsePaymentConfig 视为未配置(换算关闭) + } + return strconv.FormatFloat(*v, 'f', -1, 64) +} + func formatNonNegativeFloat(v *float64) string { if v == nil || *v < 0 { return "" diff --git a/backend/internal/service/payment_config_service_test.go b/backend/internal/service/payment_config_service_test.go index f04f4697b1..bfc69d1705 100644 --- a/backend/internal/service/payment_config_service_test.go +++ b/backend/internal/service/payment_config_service_test.go @@ -187,6 +187,23 @@ func TestParsePaymentConfig(t *testing.T) { } }) + t.Run("custom enabled types are preserved", func(t *testing.T) { + t.Parallel() + vals := map[string]string{ + SettingEnabledPaymentTypes: "alipay,ldc,usdt_trc20", + } + cfg := svc.parsePaymentConfig(vals) + want := []string{"alipay", "ldc", "usdt_trc20"} + if len(cfg.EnabledTypes) != len(want) { + t.Fatalf("EnabledTypes len = %d, want %d (%v)", len(cfg.EnabledTypes), len(want), cfg.EnabledTypes) + } + for i := range want { + if cfg.EnabledTypes[i] != want[i] { + t.Fatalf("EnabledTypes[%d] = %q, want %q (full=%v)", i, cfg.EnabledTypes[i], want[i], cfg.EnabledTypes) + } + } + }) + t.Run("empty enabled types string", func(t *testing.T) { t.Parallel() vals := map[string]string{ diff --git a/backend/internal/service/payment_fulfillment_test.go b/backend/internal/service/payment_fulfillment_test.go index b46d6a1fc8..a8c78d713c 100644 --- a/backend/internal/service/payment_fulfillment_test.go +++ b/backend/internal/service/payment_fulfillment_test.go @@ -602,8 +602,8 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { SetUserID(user.ID). SetUserEmail(user.Email). SetUserName(user.Username). - SetAmount(120). - SetPayAmount(120). + SetAmount(9.99). + SetPayAmount(71.36). SetFeeRate(0). SetRechargeCode("PAY-SUB-AFFILIATE"). SetOutTradeNo("sub2_subscription_affiliate"). @@ -636,7 +636,7 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { } settingSvc := NewSettingService(&paymentFulfillmentSettingRepoStub{values: map[string]string{ SettingKeyAffiliateEnabled: "true", - SettingKeyAffiliateRebateRate: "20", + SettingKeyAffiliateRebateRate: "15", SettingKeyAffiliateRebateFreezeHours: "0", }}, nil) subRepo := newSubscriptionUserSubRepoStub() @@ -659,7 +659,7 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { require.Len(t, affiliateRepo.accrueCalls, 1) require.Equal(t, inviterID, affiliateRepo.accrueCalls[0].inviterID) require.Equal(t, user.ID, affiliateRepo.accrueCalls[0].inviteeUserID) - require.Equal(t, 24.0, affiliateRepo.accrueCalls[0].amount) + require.InDelta(t, 1.4985, affiliateRepo.accrueCalls[0].amount, 0.00000001) require.NotNil(t, affiliateRepo.accrueCalls[0].sourceOrderID) require.Equal(t, order.ID, *affiliateRepo.accrueCalls[0].sourceOrderID) require.Equal(t, 1, subRepo.createCalls) @@ -668,8 +668,8 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("AFFILIATE_REBATE_APPLIED")). Only(ctx) require.NoError(t, err) - require.Contains(t, applied.Detail, `"baseAmount":120`) - require.Contains(t, applied.Detail, `"rebateAmount":24`) + require.Contains(t, applied.Detail, `"baseAmount":9.99`) + require.Contains(t, applied.Detail, `"rebateAmount":1.4985`) } func TestExecuteSubscriptionFulfillmentDoesNotDuplicateWorkAfterLegacySuccessAudit(t *testing.T) { diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index 29fe40b1b6..04feb8002a 100644 --- a/backend/internal/service/payment_order.go +++ b/backend/internal/service/payment_order.go @@ -16,6 +16,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/payment" "github.com/Wei-Shaw/sub2api/internal/payment/provider" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/shopspring/decimal" ) // --- Order Creation --- @@ -67,8 +68,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest return nil, err } } - // 订阅套餐 price 是直付价,余额充值倍率只影响余额充值到账,不参与订阅 pay_amount 计算。 - payAmountStr, payAmount, err := calculateCreateOrderPayAmount(limitAmount, feeRate, methodCurrency) + payAmountStr, payAmount, err := calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, methodCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err } @@ -84,7 +84,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest selectedCurrency = paymentProviderConfigCurrency(sel.ProviderKey, sel.Config) } if selectedCurrency != methodCurrency { - payAmountStr, payAmount, err = calculateCreateOrderPayAmount(limitAmount, feeRate, selectedCurrency) + payAmountStr, payAmount, err = calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, selectedCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err } @@ -453,6 +453,7 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen } return nil, classifyCreatePaymentError(req, sel.ProviderKey, err) } + sanitizeCreatePaymentResponseDetails(pr) _, err = s.entClient.PaymentOrder.UpdateOneID(order.ID). SetNillablePaymentTradeNo(psNilIfEmpty(pr.TradeNo)). SetNillablePayURL(psNilIfEmpty(pr.PayURL)). @@ -480,6 +481,22 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen return resp, nil } +func sanitizeCreatePaymentResponseDetails(pr *payment.CreatePaymentResponse) { + if pr == nil { + return + } + pr.TradeNo = removePostgresTextNUL(pr.TradeNo) + pr.PayURL = removePostgresTextNUL(pr.PayURL) + pr.QRCode = removePostgresTextNUL(pr.QRCode) +} + +func removePostgresTextNUL(value string) string { + if !strings.ContainsRune(value, 0) { + return value + } + return strings.ReplaceAll(value, "\x00", "") +} + func buildProviderCreatePaymentRequest(req CreateOrderRequest, sel *payment.InstanceSelection, orderID, amount, subject string) payment.CreatePaymentRequest { return payment.CreatePaymentRequest{ OrderID: orderID, @@ -613,6 +630,28 @@ func calculateCreateOrderPayAmount(limitAmount, feeRate float64, currency string return payAmountStr, payAmount, nil } +func calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate float64, currency, orderType string, usdToCnyRate float64) (string, float64, error) { + paymentAmount := limitAmount + if orderType == payment.OrderTypeSubscription { + paymentAmount = calculateSubscriptionGatewayBaseAmount(limitAmount, usdToCnyRate, currency) + } + return calculateCreateOrderPayAmount(paymentAmount, feeRate, currency) +} + +// calculateSubscriptionGatewayBaseAmount 计算订阅订单的网关扣款基数。 +// 换算是显式 opt-in:仅当管理员配置了订阅汇率(rate > 0,1 USD = rate CNY) +// 且网关币种为 CNY 时,按 price × rate 换算;未配置时保持 price 直付的存量行为。 +func calculateSubscriptionGatewayBaseAmount(amount, usdToCnyRate float64, currency string) float64 { + rate := normalizeSubscriptionUSDToCNYRate(usdToCnyRate) + if rate <= 0 || currency != payment.DefaultPaymentCurrency { + return amount + } + return decimal.NewFromFloat(amount). + Mul(decimal.NewFromFloat(rate)). + Round(int32(payment.CurrencyMaxFractionDigits(currency))). + InexactFloat64() +} + func validateCreateOrderAmountCurrency(amount float64, currency string) error { amountStr := strconv.FormatFloat(amount, 'f', -1, 64) if _, err := payment.AmountToMinorUnit(amountStr, currency); err != nil { diff --git a/backend/internal/service/payment_order_result_test.go b/backend/internal/service/payment_order_result_test.go index b7545ee45d..ac439ee6f2 100644 --- a/backend/internal/service/payment_order_result_test.go +++ b/backend/internal/service/payment_order_result_test.go @@ -91,6 +91,41 @@ func TestBuildCreateOrderResponseCopiesJSAPIPayload(t *testing.T) { } } +func TestSanitizeCreatePaymentResponseDetailsRemovesNULBytes(t *testing.T) { + t.Parallel() + + resp := &payment.CreatePaymentResponse{ + TradeNo: "trade\x00-no", + PayURL: "https://pay.example.com/\x00checkout", + QRCode: "wxp://payment-token\x00", + ClientSecret: "secret\x00unchanged", + } + + sanitizeCreatePaymentResponseDetails(resp) + + if strings.ContainsRune(resp.TradeNo, 0) { + t.Fatalf("trade_no still contains NUL: %q", resp.TradeNo) + } + if strings.ContainsRune(resp.PayURL, 0) { + t.Fatalf("pay_url still contains NUL: %q", resp.PayURL) + } + if strings.ContainsRune(resp.QRCode, 0) { + t.Fatalf("qr_code still contains NUL: %q", resp.QRCode) + } + if resp.TradeNo != "trade-no" { + t.Fatalf("trade_no = %q, want trade-no", resp.TradeNo) + } + if resp.PayURL != "https://pay.example.com/checkout" { + t.Fatalf("pay_url = %q, want sanitized URL", resp.PayURL) + } + if resp.QRCode != "wxp://payment-token" { + t.Fatalf("qr_code = %q, want sanitized QR code", resp.QRCode) + } + if resp.ClientSecret != "secret\x00unchanged" { + t.Fatalf("client_secret = %q, should not be touched by payment detail sanitization", resp.ClientSecret) + } +} + func TestValidateSelectedCreateOrderAmountCurrencyRejectsFractionalZeroDecimal(t *testing.T) { t.Parallel() @@ -126,27 +161,66 @@ func TestCalculateCreateOrderPayAmountUsesCurrencyPrecision(t *testing.T) { } } -func TestCalculateCreateOrderPayAmountForSubscriptionKeepsDirectPrice(t *testing.T) { +func TestCalculateCreateOrderPayAmountForSubscriptionConvertsCNYPriceWhenRateConfigured(t *testing.T) { t.Parallel() - amountStr, amount, err := calculateCreateOrderPayAmount(5, 0, "CNY") + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "CNY", payment.OrderTypeSubscription, 7.15) if err != nil { t.Fatalf("unexpected error: %v", err) } - if amountStr != "5.00" || amount != 5 { - t.Fatalf("subscription CNY pay amount = (%q, %v), want (5.00, 5)", amountStr, amount) + if amountStr != "71.43" || amount != 71.43 { + t.Fatalf("subscription CNY pay amount = (%q, %v), want (71.43, 71.43)", amountStr, amount) } } -func TestCalculateCreateOrderPayAmountForSubscriptionAppliesFeeToDirectPrice(t *testing.T) { +func TestCalculateCreateOrderPayAmountForSubscriptionAppliesFeeAfterCNYConversion(t *testing.T) { t.Parallel() - amountStr, amount, err := calculateCreateOrderPayAmount(5, 2.5, "CNY") + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 2.5, "CNY", payment.OrderTypeSubscription, 7.15) if err != nil { t.Fatalf("unexpected error: %v", err) } - if amountStr != "5.13" || amount != 5.13 { - t.Fatalf("subscription CNY pay amount with fee = (%q, %v), want (5.13, 5.13)", amountStr, amount) + if amountStr != "73.22" || amount != 73.22 { + t.Fatalf("subscription CNY pay amount with fee = (%q, %v), want (73.22, 73.22)", amountStr, amount) + } +} + +func TestCalculateCreateOrderPayAmountForSubscriptionKeepsNonCNYPrice(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "USD", payment.OrderTypeSubscription, 7.15) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "9.99" || amount != 9.99 { + t.Fatalf("subscription USD pay amount = (%q, %v), want (9.99, 9.99)", amountStr, amount) + } +} + +// 换算是 opt-in:未配置汇率(rate=0)时,CNY 订阅保持 price 直付的存量行为。 +// 该测试锁住存量部署升级后行为不变的兼容承诺。 +func TestCalculateCreateOrderPayAmountForSubscriptionKeepsDirectPriceWhenRateDisabled(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "CNY", payment.OrderTypeSubscription, 0) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "9.99" || amount != 9.99 { + t.Fatalf("subscription CNY pay amount without rate = (%q, %v), want (9.99, 9.99)", amountStr, amount) + } +} + +// 汇率只作用于订阅订单,余额充值订单不受影响。 +func TestCalculateCreateOrderPayAmountForBalanceIgnoresSubscriptionRate(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(50, 0, "CNY", payment.OrderTypeBalance, 7.15) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "50.00" || amount != 50 { + t.Fatalf("balance CNY pay amount = (%q, %v), want (50.00, 50)", amountStr, amount) } } diff --git a/backend/internal/service/payment_resume_service_test.go b/backend/internal/service/payment_resume_service_test.go index 7e0adc2de8..17b637fa23 100644 --- a/backend/internal/service/payment_resume_service_test.go +++ b/backend/internal/service/payment_resume_service_test.go @@ -26,9 +26,10 @@ func TestNormalizeVisibleMethods(t *testing.T) { " wxpay_direct ", "wxpay", "stripe", + "ldc", }) - want := []string{"alipay", "wxpay", "stripe"} + want := []string{"alipay", "wxpay", "stripe", "ldc"} if len(got) != len(want) { t.Fatalf("NormalizeVisibleMethods len = %d, want %d (%v)", len(got), len(want), got) } @@ -39,6 +40,21 @@ func TestNormalizeVisibleMethods(t *testing.T) { } } +func TestEnabledVisibleMethodsForEasyPayIncludesCustomSupportedTypes(t *testing.T) { + t.Parallel() + + got := enabledVisibleMethodsForProvider(payment.TypeEasyPay, "alipay,ldc,usdt_trc20") + want := []string{"alipay", "ldc", "usdt_trc20"} + if len(got) != len(want) { + t.Fatalf("enabledVisibleMethodsForProvider len = %d, want %d (%v)", len(got), len(want), got) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("enabledVisibleMethodsForProvider[%d] = %q, want %q (full=%v)", i, got[i], want[i], got) + } + } +} + func TestNormalizePaymentSource(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_visible_method_instances.go b/backend/internal/service/payment_visible_method_instances.go index 899bd7a020..97b3b1ef66 100644 --- a/backend/internal/service/payment_visible_method_instances.go +++ b/backend/internal/service/payment_visible_method_instances.go @@ -16,8 +16,7 @@ func enabledVisibleMethodsForProvider(providerKey, supportedTypes string) []stri methodSet := make(map[string]struct{}, 2) addMethod := func(method string) { method = NormalizeVisibleMethod(method) - switch method { - case payment.TypeAlipay, payment.TypeWxpay: + if method != "" { methodSet[method] = struct{}{} } } @@ -55,6 +54,14 @@ func enabledVisibleMethodsForProvider(providerKey, supportedTypes string) []stri for _, method := range []string{payment.TypeAlipay, payment.TypeWxpay} { if _, ok := methodSet[method]; ok { methods = append(methods, method) + delete(methodSet, method) + } + } + for _, supportedType := range splitTypes(supportedTypes) { + method := NormalizeVisibleMethod(supportedType) + if _, ok := methodSet[method]; ok { + methods = append(methods, method) + delete(methodSet, method) } } return methods @@ -215,7 +222,7 @@ func (s *PaymentConfigService) resolveEnabledVisibleMethodInstance( } method = NormalizeVisibleMethod(method) - if method != payment.TypeAlipay && method != payment.TypeWxpay { + if method == "" { return nil, nil } diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index bd0c30df45..1a0b603169 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -798,6 +798,13 @@ func (s *PricingService) matchOpenAIModel(model string) *LiteLLMModelPricing { } } + // GPT-5.6(sol / terra / luna)回退到 GPT-5.4 定价 + if strings.HasPrefix(model, "gpt-5.6") { + logger.With(zap.String("component", "service.pricing")). + Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.4(static)")) + return openAIGPT54FallbackPricing + } + // GPT-5.5 回退到 GPT-5.4 定价 if strings.HasPrefix(model, "gpt-5.5") { logger.With(zap.String("component", "service.pricing")). diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 9e968853f4..50d38e7d13 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -116,6 +116,14 @@ func (s *RateLimitService) SetAccountRuntimeBlocker(blocker AccountRuntimeBlocke s.runtimeBlocker = blocker } +func (s *RateLimitService) IsOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx context.Context) bool { + if s == nil || s.settingService == nil { + return false + } + gateway := &OpenAIGatewayService{rateLimitService: s} + return gateway.isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx) +} + func (s *RateLimitService) notifyAccountSchedulingBlocked(account *Account, until time.Time, reason string) { if s == nil || s.runtimeBlocker == nil || account == nil { return @@ -186,9 +194,14 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc // otherwise a broad "rate limit" keyword rule can shorten a multi-hour // cooldown to a local temporary pause. if statusCode == http.StatusTooManyRequests && account.Platform == PlatformAnthropic { + // 7d_oi 是 Fable 模型专属的 7d 窗口:只标记模型级限流,账号对其他模型仍可调度。 + fableLimited := s.persistAnthropicFableWindowLimit(ctx, account, headers) if s.persistAnthropicExhaustedWindowLimit(ctx, account, headers) { return false } + if fableLimited { + return false + } } // 先尝试临时不可调度规则(401除外) @@ -287,6 +300,20 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc if upstreamMsg != "" { msg = "OAuth 401: " + upstreamMsg } + if authAccount.Platform == PlatformAntigravity { + extraUpdates := antigravityForceTokenRefreshExtra("401_invalid") + if err := s.accountRepo.UpdateExtra(ctx, authAccount.ID, extraUpdates); err != nil { + slog.Warn("antigravity_401_force_refresh_mark_failed", "account_id", authAccount.ID, "error", err) + } else { + if authAccount.Extra == nil { + authAccount.Extra = make(map[string]any, len(extraUpdates)) + } + for k, v := range extraUpdates { + authAccount.Extra[k] = v + } + slog.Info("antigravity_401_force_refresh_marked", "account_id", authAccount.ID) + } + } cooldownMinutes := s.cfg.RateLimit.OAuth401CooldownMinutes if cooldownMinutes <= 0 { cooldownMinutes = 10 @@ -1140,11 +1167,25 @@ func selectAnthropicExhaustedWindow(headers http.Header, now time.Time) *anthrop } func isAnthropic5hRejected(headers http.Header) bool { - return strings.EqualFold(strings.TrimSpace(headers.Get("anthropic-ratelimit-unified-5h-status")), "rejected") + return isAnthropicWindowRejected(headers, "5h") +} + +func isAnthropicWindowRejected(headers http.Header, window string) bool { + return strings.EqualFold(strings.TrimSpace(headers.Get("anthropic-ratelimit-unified-"+window+"-status")), "rejected") } func parseAnthropicWindowReset(headers http.Header, window string, now time.Time) (time.Time, bool) { - raw := strings.TrimSpace(headers.Get("anthropic-ratelimit-unified-" + window + "-reset")) + maxAge := 8 * 24 * time.Hour + if window == "5h" { + maxAge = 6 * time.Hour + } + return parseAnthropicResetTimestamp(headers.Get("anthropic-ratelimit-unified-"+window+"-reset"), now, maxAge) +} + +// parseAnthropicResetTimestamp 解析 Anthropic reset 头的 Unix 时间戳(自动识别毫秒), +// 并校验落在 (now, now+maxAge] 的合理区间内。 +func parseAnthropicResetTimestamp(raw string, now time.Time, maxAge time.Duration) (time.Time, bool) { + raw = strings.TrimSpace(raw) if raw == "" { return time.Time{}, false } @@ -1156,15 +1197,7 @@ func parseAnthropicWindowReset(headers http.Header, window string, now time.Time ts = ts / 1000 } resetAt := time.Unix(ts, 0) - if !resetAt.After(now) { - return time.Time{}, false - } - - maxAge := 8 * 24 * time.Hour - if window == "5h" { - maxAge = 6 * time.Hour - } - if resetAt.After(now.Add(maxAge)) { + if !resetAt.After(now) || resetAt.After(now.Add(maxAge)) { return time.Time{}, false } return resetAt, true @@ -1218,6 +1251,76 @@ func (s *RateLimitService) persistAnthropicExhaustedWindowLimit(ctx context.Cont return true } +const anthropicFableWindowReason = "anthropic_7d_oi_window_exhausted" + +// selectAnthropicFableWindowLimit parses the Anthropic 7d_oi per-model window +// headers (the Fable-only 7d window, e.g. anthropic-ratelimit-unified-7d_oi-*). +// Unlike 5h/7d, exhaustion of this window only limits the Fable model family — +// the account must stay schedulable for other models. +// +// The 7d_oi surpassed-threshold header carries a float ("1.0") rather than +// "true", so exhaustion is detected via status=rejected or utilization >= 1.0. +// When the 7d_oi reset header is missing, the aggregated +// anthropic-ratelimit-unified-reset is used (it mirrors the binding claim's +// reset when 7d_oi is the representative claim). +func selectAnthropicFableWindowLimit(headers http.Header, now time.Time) *anthropicWindowLimit { + if !isAnthropicWindowRejected(headers, "7d_oi") && !isAnthropicWindowExceeded(headers, "7d_oi") { + return nil + } + resetAt, ok := parseAnthropicWindowReset(headers, "7d_oi", now) + if !ok { + resetAt, ok = parseAnthropicAggregateReset(headers, now) + } + if !ok { + return nil + } + return &anthropicWindowLimit{ + window: "7d_oi", + resetAt: resetAt, + reason: anthropicFableWindowReason, + } +} + +// parseAnthropicAggregateReset parses the aggregated +// anthropic-ratelimit-unified-reset header with the same sanity checks as the +// per-window variant (7d scale). +func parseAnthropicAggregateReset(headers http.Header, now time.Time) (time.Time, bool) { + return parseAnthropicResetTimestamp(headers.Get("anthropic-ratelimit-unified-reset"), now, 8*24*time.Hour) +} + +// persistAnthropicFableWindowLimit marks the Fable model family as rate limited +// when the 7d_oi window is exhausted. Returns true when the 7d_oi window was the +// (or a) trigger of this 429, so the caller must not fall through to logic that +// would mark the whole account as rate limited. +func (s *RateLimitService) persistAnthropicFableWindowLimit(ctx context.Context, account *Account, headers http.Header) bool { + if s == nil || s.accountRepo == nil || account == nil { + return false + } + now := time.Now() + limit := selectAnthropicFableWindowLimit(headers, now) + if limit == nil { + return false + } + // 429 响应头本身携带最新的窗口用量(7d_oi utilization=1.0)。限流期内 + // Fable 请求不再调度到该账号,若不在此处采样,7d F 进度条会冻结在 + // 限流前的旧值直到窗口重置。 + s.samplePassiveUsageFromHeaders(ctx, account, headers) + if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, anthropicFableRateLimitKey, limit.resetAt, limit.reason); err != nil { + slog.Warn("anthropic_fable_window_rate_limit_set_failed", + "account_id", account.ID, + "scope", anthropicFableRateLimitKey, + "reset_at", limit.resetAt, + "error", err) + return true + } + slog.Info("anthropic_fable_window_model_rate_limited", + "account_id", account.ID, + "scope", anthropicFableRateLimitKey, + "reset_at", limit.resetAt, + "reset_in", time.Until(limit.resetAt).Truncate(time.Second)) + return true +} + // calculateAnthropic429ResetTime parses Anthropic's per-window rate-limit headers // to determine which window (5h or 7d) actually triggered the 429. // @@ -1541,10 +1644,12 @@ func (s *RateLimitService) UpdateSessionWindow(ctx context.Context, account *Acc // 窗口重置时清除旧的 utilization 和被动采样数据,避免残留上个窗口的数据 if windowEnd != nil && needInitWindow { _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ - "session_window_utilization": nil, - "passive_usage_7d_utilization": nil, - "passive_usage_7d_reset": nil, - "passive_usage_sampled_at": nil, + "session_window_utilization": nil, + "passive_usage_7d_utilization": nil, + "passive_usage_7d_reset": nil, + "passive_usage_7d_oi_utilization": nil, + "passive_usage_7d_oi_reset": nil, + "passive_usage_sampled_at": nil, }) } @@ -1552,8 +1657,21 @@ func (s *RateLimitService) UpdateSessionWindow(ctx context.Context, account *Acc slog.Warn("session_window_update_failed", "account_id", account.ID, "error", err) } - // 被动采样:从响应头收集 5h + 7d utilization,合并为一次 DB 写入 - extraUpdates := make(map[string]any, 4) + // 被动采样:从响应头收集 5h + 7d + 7d_oi utilization,合并为一次 DB 写入 + s.samplePassiveUsageFromHeaders(ctx, account, headers) + + // 如果状态为allowed且之前有限流,说明窗口已重置,清除限流状态 + if status == "allowed" && account.IsRateLimited() { + if err := s.ClearRateLimit(ctx, account.ID); err != nil { + slog.Warn("rate_limit_clear_failed", "account_id", account.ID, "error", err) + } + } +} + +// samplePassiveUsageFromHeaders 从 Anthropic 响应头收集 5h/7d/7d_oi 的 +// utilization 与 reset 被动采样数据,合并为一次 Extra 写入。无数据时不写。 +func (s *RateLimitService) samplePassiveUsageFromHeaders(ctx context.Context, account *Account, headers http.Header) { + extraUpdates := make(map[string]any, 6) // 5h utilization(0-1 小数),供 estimateSetupTokenUsage 使用 if utilStr := headers.Get("anthropic-ratelimit-unified-5h-utilization"); utilStr != "" { if util, err := strconv.ParseFloat(utilStr, 64); err == nil { @@ -1575,19 +1693,27 @@ func (s *RateLimitService) UpdateSessionWindow(ctx context.Context, account *Acc extraUpdates["passive_usage_7d_reset"] = ts } } + // 7d_oi (Fable 专属 7d 窗口) utilization(0-1 小数) + if utilStr := headers.Get("anthropic-ratelimit-unified-7d_oi-utilization"); utilStr != "" { + if util, err := strconv.ParseFloat(utilStr, 64); err == nil { + extraUpdates["passive_usage_7d_oi_utilization"] = util + } + } + // 7d_oi reset timestamp + if resetStr := headers.Get("anthropic-ratelimit-unified-7d_oi-reset"); resetStr != "" { + if ts, err := strconv.ParseInt(resetStr, 10, 64); err == nil { + if ts > 1e11 { + ts = ts / 1000 + } + extraUpdates["passive_usage_7d_oi_reset"] = ts + } + } if len(extraUpdates) > 0 { extraUpdates["passive_usage_sampled_at"] = time.Now().UTC().Format(time.RFC3339) if err := s.accountRepo.UpdateExtra(ctx, account.ID, extraUpdates); err != nil { slog.Warn("passive_usage_update_failed", "account_id", account.ID, "error", err) } } - - // 如果状态为allowed且之前有限流,说明窗口已重置,清除限流状态 - if status == "allowed" && account.IsRateLimited() { - if err := s.ClearRateLimit(ctx, account.ID); err != nil { - slog.Warn("rate_limit_clear_failed", "account_id", account.ID, "error", err) - } - } } // ClearRateLimit 清除账号的限流状态 diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go index 09afb9314b..48e6a41def 100644 --- a/backend/internal/service/ratelimit_service_401_test.go +++ b/backend/internal/service/ratelimit_service_401_test.go @@ -18,7 +18,9 @@ type rateLimitAccountRepoStub struct { setErrorCalls int tempCalls int updateCredentialsCalls int + updateExtraCalls int lastCredentials map[string]any + lastExtraUpdates map[string]any lastErrorMsg string lastTempReason string lastErrorID int64 @@ -45,6 +47,12 @@ func (r *rateLimitAccountRepoStub) UpdateCredentials(ctx context.Context, id int return nil } +func (r *rateLimitAccountRepoStub) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error { + r.updateExtraCalls++ + r.lastExtraUpdates = shallowCopyMap(updates) + return nil +} + type tokenCacheInvalidatorRecorder struct { accounts []*Account err error @@ -133,6 +141,10 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t require.Equal(t, 1, repo.tempCalls) require.Equal(t, int64(100), repo.lastTempID) require.Contains(t, repo.lastTempReason, "invalid or expired credentials") + require.Equal(t, 1, repo.updateExtraCalls) + require.Equal(t, true, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey]) + require.Equal(t, "401_invalid", repo.lastExtraUpdates[antigravityForceTokenRefreshReasonExtraKey]) + require.Equal(t, true, account.Extra[antigravityForceTokenRefreshExtraKey]) require.Len(t, invalidator.accounts, 1) require.Equal(t, int64(100), invalidator.accounts[0].ID) }) @@ -245,6 +257,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredential require.True(t, shouldDisable) require.Equal(t, 0, repo.updateCredentialsCalls, "401 handler must not write credentials back from the request-start snapshot") + require.Equal(t, 0, repo.updateExtraCalls, "OpenAI 401 must not set Antigravity force-refresh marker") require.Equal(t, 1, repo.tempCalls, "401 handler should still set temp-unschedulable cooldown") require.Nil(t, repo.lastCredentials, "no credentials should have been persisted") } diff --git a/backend/internal/service/ratelimit_service_anthropic_test.go b/backend/internal/service/ratelimit_service_anthropic_test.go index eaeaf30e60..0e75b4f914 100644 --- a/backend/internal/service/ratelimit_service_anthropic_test.go +++ b/backend/internal/service/ratelimit_service_anthropic_test.go @@ -2,6 +2,7 @@ package service import ( "net/http" + "strconv" "testing" "time" ) @@ -181,6 +182,142 @@ func TestIsAnthropicWindowExceeded(t *testing.T) { } } +func TestSelectAnthropicFableWindowLimit_RejectedStatus(t *testing.T) { + now := time.Now() + reset := now.Add(80 * time.Hour).Truncate(time.Second) + + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "rejected") + headers.Set("anthropic-ratelimit-unified-7d_oi-utilization", "1.0") + headers.Set("anthropic-ratelimit-unified-7d_oi-surpassed-threshold", "1.0") + headers.Set("anthropic-ratelimit-unified-7d_oi-reset", strconv.FormatInt(reset.Unix(), 10)) + + limit := selectAnthropicFableWindowLimit(headers, now) + if limit == nil { + t.Fatal("expected non-nil limit") + } + if !limit.resetAt.Equal(reset) { + t.Errorf("expected resetAt=%v, got %v", reset, limit.resetAt) + } + if limit.reason != anthropicFableWindowReason { + t.Errorf("expected reason=%q, got %q", anthropicFableWindowReason, limit.reason) + } +} + +func TestSelectAnthropicFableWindowLimit_UtilizationOnly(t *testing.T) { + // 无 status 头时,utilization >= 1.0 也应视为超限 + now := time.Now() + reset := now.Add(3 * 24 * time.Hour).Truncate(time.Second) + + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-utilization", "1.0") + headers.Set("anthropic-ratelimit-unified-7d_oi-reset", strconv.FormatInt(reset.Unix(), 10)) + + limit := selectAnthropicFableWindowLimit(headers, now) + if limit == nil { + t.Fatal("expected non-nil limit") + } + if !limit.resetAt.Equal(reset) { + t.Errorf("expected resetAt=%v, got %v", reset, limit.resetAt) + } +} + +func TestSelectAnthropicFableWindowLimit_AllowedReturnsNil(t *testing.T) { + now := time.Now() + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "allowed") + headers.Set("anthropic-ratelimit-unified-7d_oi-utilization", "0.56") + headers.Set("anthropic-ratelimit-unified-7d_oi-reset", strconv.FormatInt(now.Add(80*time.Hour).Unix(), 10)) + + if limit := selectAnthropicFableWindowLimit(headers, now); limit != nil { + t.Errorf("expected nil limit for allowed window, got %+v", limit) + } +} + +func TestSelectAnthropicFableWindowLimit_NoHeadersReturnsNil(t *testing.T) { + if limit := selectAnthropicFableWindowLimit(http.Header{}, time.Now()); limit != nil { + t.Errorf("expected nil limit for empty headers, got %+v", limit) + } +} + +func TestSelectAnthropicFableWindowLimit_FallsBackToAggregateReset(t *testing.T) { + // 7d_oi-reset 缺失时回退聚合 anthropic-ratelimit-unified-reset + now := time.Now() + reset := now.Add(80 * time.Hour).Truncate(time.Second) + + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "rejected") + headers.Set("anthropic-ratelimit-unified-reset", strconv.FormatInt(reset.Unix(), 10)) + + limit := selectAnthropicFableWindowLimit(headers, now) + if limit == nil { + t.Fatal("expected non-nil limit via aggregate reset fallback") + } + if !limit.resetAt.Equal(reset) { + t.Errorf("expected resetAt=%v, got %v", reset, limit.resetAt) + } +} + +func TestSelectAnthropicFableWindowLimit_RejectedWithoutAnyResetReturnsNil(t *testing.T) { + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "rejected") + + if limit := selectAnthropicFableWindowLimit(headers, time.Now()); limit != nil { + t.Errorf("expected nil limit when no reset time available, got %+v", limit) + } +} + +func TestParseAnthropicAggregateReset(t *testing.T) { + now := time.Now() + future := now.Add(80 * time.Hour).Truncate(time.Second) + + tests := []struct { + name string + value string + want time.Time + wantOK bool + }{ + {"valid seconds", strconv.FormatInt(future.Unix(), 10), future, true}, + {"valid milliseconds", strconv.FormatInt(future.UnixMilli(), 10), future, true}, + {"empty", "", time.Time{}, false}, + {"garbage", "abc", time.Time{}, false}, + {"in the past", strconv.FormatInt(now.Add(-time.Hour).Unix(), 10), time.Time{}, false}, + {"too far in the future", strconv.FormatInt(now.Add(30*24*time.Hour).Unix(), 10), time.Time{}, false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + headers := http.Header{} + if tc.value != "" { + headers.Set("anthropic-ratelimit-unified-reset", tc.value) + } + got, ok := parseAnthropicAggregateReset(headers, now) + if ok != tc.wantOK { + t.Fatalf("expected ok=%v, got %v", tc.wantOK, ok) + } + if ok && !got.Equal(tc.want) { + t.Errorf("expected %v, got %v", tc.want, got) + } + }) + } +} + +func TestIsAnthropicWindowRejected(t *testing.T) { + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "Rejected") + headers.Set("anthropic-ratelimit-unified-5h-status", "allowed") + + if !isAnthropicWindowRejected(headers, "7d_oi") { + t.Error("expected 7d_oi to be rejected (case insensitive)") + } + if isAnthropicWindowRejected(headers, "5h") { + t.Error("expected 5h not rejected") + } + if isAnthropicWindowRejected(headers, "7d") { + t.Error("expected missing 7d status not rejected") + } +} + // assertAnthropicResult is a test helper that verifies the result is non-nil and // has the expected resetAt unix timestamp. func assertAnthropicResult(t *testing.T, result *anthropic429Result, wantUnix int64) { diff --git a/backend/internal/service/ratelimit_service_anthropic_window_limit_test.go b/backend/internal/service/ratelimit_service_anthropic_window_limit_test.go index 6168b250a6..a13e8fe333 100644 --- a/backend/internal/service/ratelimit_service_anthropic_window_limit_test.go +++ b/backend/internal/service/ratelimit_service_anthropic_window_limit_test.go @@ -14,9 +14,14 @@ import ( type anthropicWindowLimitRepo struct { mockAccountRepoForGemini - rateLimitCalls int - tempUnschedCalls int - lastRateLimitReset time.Time + rateLimitCalls int + tempUnschedCalls int + lastRateLimitReset time.Time + modelRateLimitCalls int + lastModelRateLimitScope string + lastModelRateLimitReset time.Time + sessionWindowCalls int + lastExtraUpdates map[string]any } func (r *anthropicWindowLimitRepo) SetRateLimited(_ context.Context, _ int64, resetAt time.Time) error { @@ -30,6 +35,23 @@ func (r *anthropicWindowLimitRepo) SetTempUnschedulable(_ context.Context, _ int return nil } +func (r *anthropicWindowLimitRepo) SetModelRateLimit(_ context.Context, _ int64, scope string, resetAt time.Time, _ ...string) error { + r.modelRateLimitCalls++ + r.lastModelRateLimitScope = scope + r.lastModelRateLimitReset = resetAt + return nil +} + +func (r *anthropicWindowLimitRepo) UpdateSessionWindow(_ context.Context, _ int64, _, _ *time.Time, _ string) error { + r.sessionWindowCalls++ + return nil +} + +func (r *anthropicWindowLimitRepo) UpdateExtra(_ context.Context, _ int64, updates map[string]any) error { + r.lastExtraUpdates = updates + return nil +} + func TestHandleUpstreamError_AnthropicWindowLimitPreemptsTempUnschedRule(t *testing.T) { resetAt := time.Now().Add(3 * time.Hour).Truncate(time.Second) headers := http.Header{} @@ -66,3 +88,142 @@ func TestHandleUpstreamError_AnthropicWindowLimitPreemptsTempUnschedRule(t *test require.Equal(t, 1, repo.rateLimitCalls) require.Equal(t, resetAt, repo.lastRateLimitReset) } + +// fable429Headers 构造 7d_oi(Fable 专属 7d 窗口)触发 429 的完整响应头, +// 数值取自真实抓包(5h/7d 均 allowed,仅 7d_oi rejected)。 +func fable429Headers(reset5h, resetOI time.Time) http.Header { + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-5h-reset", strconv.FormatInt(reset5h.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-5h-status", "allowed") + headers.Set("anthropic-ratelimit-unified-5h-utilization", "0.41") + headers.Set("anthropic-ratelimit-unified-7d-reset", strconv.FormatInt(resetOI.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-7d-status", "allowed") + headers.Set("anthropic-ratelimit-unified-7d-utilization", "0.56") + headers.Set("anthropic-ratelimit-unified-7d_oi-reset", strconv.FormatInt(resetOI.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-7d_oi-status", "rejected") + headers.Set("anthropic-ratelimit-unified-7d_oi-surpassed-threshold", "1.0") + headers.Set("anthropic-ratelimit-unified-7d_oi-utilization", "1.0") + headers.Set("anthropic-ratelimit-unified-fallback-percentage", "0.5") + headers.Set("anthropic-ratelimit-unified-overage-disabled-reason", "org_level_disabled") + headers.Set("anthropic-ratelimit-unified-overage-status", "rejected") + headers.Set("anthropic-ratelimit-unified-representative-claim", "seven_day_overage_included") + headers.Set("anthropic-ratelimit-unified-reset", strconv.FormatInt(resetOI.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-status", "rejected") + return headers +} + +func TestHandleUpstreamError_Anthropic7dOiOnlyMarksModelRateLimit(t *testing.T) { + now := time.Now() + reset5h := now.Add(2 * time.Hour).Truncate(time.Second) + resetOI := now.Add(80 * time.Hour).Truncate(time.Second) + headers := fable429Headers(reset5h, resetOI) + + repo := &anthropicWindowLimitRepo{} + svc := NewRateLimitService(repo, nil, nil, nil, nil) + account := &Account{ + ID: 42, + Type: AccountTypeOAuth, + Platform: PlatformAnthropic, + Credentials: map[string]any{ + "temp_unschedulable_enabled": true, + "temp_unschedulable_rules": []any{ + map[string]any{ + "error_code": float64(http.StatusTooManyRequests), + "keywords": []any{"rate limit"}, + "duration_minutes": float64(10), + }, + }, + }, + } + + shouldDisable := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusTooManyRequests, + headers, + []byte(`{"type":"error","error":{"type":"rate_limit_error","message":"This request would exceed your account's rate limit. Please try again later."}}`), + "claude-fable-5", + ) + + require.False(t, shouldDisable) + require.Zero(t, repo.rateLimitCalls, "7d_oi (Fable-only) window must not mark the whole account rate limited") + require.Zero(t, repo.tempUnschedCalls, "7d_oi window must not trigger local temp-unsched rules") + require.Zero(t, repo.sessionWindowCalls, "7d_oi window must not rewrite the 5h session window as rejected") + require.Equal(t, 1, repo.modelRateLimitCalls) + require.Equal(t, anthropicFableRateLimitKey, repo.lastModelRateLimitScope) + require.Equal(t, resetOI, repo.lastModelRateLimitReset) + + // 429 响应头也要被动采样,避免 7d F 进度条在限流期内冻结在旧值 + require.NotNil(t, repo.lastExtraUpdates) + require.Equal(t, 1.0, repo.lastExtraUpdates["passive_usage_7d_oi_utilization"]) + require.Equal(t, resetOI.Unix(), repo.lastExtraUpdates["passive_usage_7d_oi_reset"]) + require.Equal(t, 0.41, repo.lastExtraUpdates["session_window_utilization"]) +} + +func TestHandleUpstreamError_Anthropic5hWindowStillWinsOver7dOi(t *testing.T) { + // 5h 窗口 rejected 时必须仍按账号级限流处理(用 5h reset),同时记录 Fable 模型限流。 + now := time.Now() + reset5h := now.Add(2 * time.Hour).Truncate(time.Second) + resetOI := now.Add(80 * time.Hour).Truncate(time.Second) + headers := fable429Headers(reset5h, resetOI) + headers.Set("anthropic-ratelimit-unified-5h-status", "rejected") + headers.Set("anthropic-ratelimit-unified-5h-utilization", "1.0") + + repo := &anthropicWindowLimitRepo{} + svc := NewRateLimitService(repo, nil, nil, nil, nil) + account := &Account{ID: 42, Type: AccountTypeOAuth, Platform: PlatformAnthropic} + + svc.HandleUpstreamError(context.Background(), account, http.StatusTooManyRequests, headers, nil, "claude-fable-5") + + require.Equal(t, 1, repo.rateLimitCalls, "exhausted 5h window must still rate limit the account") + require.Equal(t, reset5h, repo.lastRateLimitReset) + require.Equal(t, 1, repo.modelRateLimitCalls) + require.Equal(t, anthropicFableRateLimitKey, repo.lastModelRateLimitScope) +} + +func TestHandleUpstreamError_AnthropicAccountWindowStillWinsOver7dOi(t *testing.T) { + // 7d 窗口真超限时必须仍按账号级限流处理,同时记录 Fable 模型限流。 + now := time.Now() + reset5h := now.Add(2 * time.Hour).Truncate(time.Second) + resetOI := now.Add(80 * time.Hour).Truncate(time.Second) + headers := fable429Headers(reset5h, resetOI) + headers.Set("anthropic-ratelimit-unified-7d-status", "rejected") + headers.Set("anthropic-ratelimit-unified-7d-utilization", "1.02") + + repo := &anthropicWindowLimitRepo{} + svc := NewRateLimitService(repo, nil, nil, nil, nil) + account := &Account{ID: 42, Type: AccountTypeOAuth, Platform: PlatformAnthropic} + + svc.HandleUpstreamError(context.Background(), account, http.StatusTooManyRequests, headers, nil, "claude-fable-5") + + require.Equal(t, 1, repo.rateLimitCalls, "exhausted 7d window must still rate limit the account") + require.Equal(t, resetOI, repo.lastRateLimitReset) + require.Equal(t, 1, repo.modelRateLimitCalls, "Fable model rate limit should also be recorded") + require.Equal(t, anthropicFableRateLimitKey, repo.lastModelRateLimitScope) +} + +func TestHandleUpstreamError_Anthropic429Without7dOiKeepsLegacyBehavior(t *testing.T) { + // 无 7d_oi 头、5h/7d 均未超限的 429:保持旧行为(按较早 reset 标记账号限流)。 + now := time.Now() + reset5h := now.Add(2 * time.Hour).Truncate(time.Second) + reset7d := now.Add(80 * time.Hour).Truncate(time.Second) + + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-5h-reset", strconv.FormatInt(reset5h.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-5h-status", "allowed") + headers.Set("anthropic-ratelimit-unified-5h-utilization", "0.41") + headers.Set("anthropic-ratelimit-unified-7d-reset", strconv.FormatInt(reset7d.Unix(), 10)) + headers.Set("anthropic-ratelimit-unified-7d-status", "allowed") + headers.Set("anthropic-ratelimit-unified-7d-utilization", "0.56") + + repo := &anthropicWindowLimitRepo{} + svc := NewRateLimitService(repo, nil, nil, nil, nil) + account := &Account{ID: 42, Type: AccountTypeOAuth, Platform: PlatformAnthropic} + + svc.HandleUpstreamError(context.Background(), account, http.StatusTooManyRequests, headers, nil, "claude-fable-5") + + require.Zero(t, repo.modelRateLimitCalls, "no 7d_oi signal → no model rate limit") + require.Equal(t, 1, repo.rateLimitCalls) + require.Equal(t, reset5h, repo.lastRateLimitReset, "legacy path picks the sooner reset") + require.Equal(t, 1, repo.sessionWindowCalls) +} diff --git a/backend/internal/service/ratelimit_session_window_test.go b/backend/internal/service/ratelimit_session_window_test.go index 9337ae5f8f..279a31ccdc 100644 --- a/backend/internal/service/ratelimit_session_window_test.go +++ b/backend/internal/service/ratelimit_session_window_test.go @@ -87,6 +87,9 @@ func (m *sessionWindowMockRepo) List(context.Context, pagination.PaginationParam func (m *sessionWindowMockRepo) ListWithFilters(context.Context, pagination.PaginationParams, string, string, string, string, int64, string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected") } +func (m *sessionWindowMockRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + panic("unexpected") +} func (m *sessionWindowMockRepo) ListByGroup(context.Context, int64) ([]Account, error) { panic("unexpected") } @@ -367,6 +370,59 @@ func TestUpdateSessionWindow_NoClearUtilizationOnCorrection(t *testing.T) { } } +func TestUpdateSessionWindow_SamplesFable7dOiHeaders(t *testing.T) { + // 被动采样应收集 7d_oi(Fable 专属 7d 窗口)的 utilization 和 reset。 + existingEnd := time.Now().Add(3 * time.Hour) + resetOIUnix := time.Now().Add(80 * time.Hour).Unix() + + repo := &sessionWindowMockRepo{} + svc := newRateLimitServiceForTest(repo) + + account := &Account{ID: 90, SessionWindowEnd: &existingEnd} // needInitWindow=false + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-5h-status", "allowed") + headers.Set("anthropic-ratelimit-unified-7d_oi-utilization", "0.87") + headers.Set("anthropic-ratelimit-unified-7d_oi-reset", fmt.Sprintf("%d", resetOIUnix)) + + svc.UpdateSessionWindow(context.Background(), account, headers) + + if len(repo.updateExtraCalls) != 1 { + t.Fatalf("expected 1 UpdateExtra call, got %d", len(repo.updateExtraCalls)) + } + updates := repo.updateExtraCalls[0].Updates + if val, ok := updates["passive_usage_7d_oi_utilization"].(float64); !ok || val != 0.87 { + t.Errorf("expected passive_usage_7d_oi_utilization=0.87, got %v", updates["passive_usage_7d_oi_utilization"]) + } + if val, ok := updates["passive_usage_7d_oi_reset"].(int64); !ok || val != resetOIUnix { + t.Errorf("expected passive_usage_7d_oi_reset=%d, got %v", resetOIUnix, updates["passive_usage_7d_oi_reset"]) + } +} + +func TestUpdateSessionWindow_ClearsFable7dOiOnWindowReset(t *testing.T) { + // 5h 窗口重置时应连同清除 7d_oi 被动采样数据,与 7d 行为一致。 + resetUnix := time.Now().Add(3 * time.Hour).Unix() + + repo := &sessionWindowMockRepo{} + svc := newRateLimitServiceForTest(repo) + + account := &Account{ID: 91} // no existing window → needInitWindow=true + headers := http.Header{} + headers.Set("anthropic-ratelimit-unified-5h-status", "allowed") + headers.Set("anthropic-ratelimit-unified-5h-reset", fmt.Sprintf("%d", resetUnix)) + + svc.UpdateSessionWindow(context.Background(), account, headers) + + if len(repo.updateExtraCalls) != 1 { + t.Fatalf("expected 1 UpdateExtra (clear) call, got %d", len(repo.updateExtraCalls)) + } + clearUpdates := repo.updateExtraCalls[0].Updates + for _, key := range []string{"passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset"} { + if val, present := clearUpdates[key]; !present || val != nil { + t.Errorf("expected %s cleared to nil on window reset, got present=%v val=%v", key, present, val) + } + } +} + func TestUpdateSessionWindow_NoStatusHeader(t *testing.T) { // Should return immediately if no status header. repo := &sessionWindowMockRepo{} diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 1a5e677a12..3aeb611418 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -1921,6 +1921,9 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting if err != nil { return nil, err } + if err := s.normalizeOpenAIAdvancedSchedulerOverrides(settings); err != nil { + return nil, err + } settings.PaymentVisibleMethodAlipaySource = alipaySource settings.PaymentVisibleMethodWxpaySource = wxpaySource settings.WeChatConnectAppID = strings.TrimSpace(settings.WeChatConnectAppID) @@ -2229,6 +2232,18 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting updates[SettingPaymentVisibleMethodAlipayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodAlipayEnabled) updates[SettingPaymentVisibleMethodWxpayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodWxpayEnabled) updates[openAIAdvancedSchedulerSettingKey] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerEnabled) + updates[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerStickyWeightedEnabled) + updates[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled) + updates[SettingKeyOpenAIAdvancedSchedulerLBTopK] = settings.OpenAIAdvancedSchedulerLBTopK + updates[SettingKeyOpenAIAdvancedSchedulerWeightPriority] = settings.OpenAIAdvancedSchedulerWeightPriority + updates[SettingKeyOpenAIAdvancedSchedulerWeightLoad] = settings.OpenAIAdvancedSchedulerWeightLoad + updates[SettingKeyOpenAIAdvancedSchedulerWeightQueue] = settings.OpenAIAdvancedSchedulerWeightQueue + updates[SettingKeyOpenAIAdvancedSchedulerWeightErrorRate] = settings.OpenAIAdvancedSchedulerWeightErrorRate + updates[SettingKeyOpenAIAdvancedSchedulerWeightTTFT] = settings.OpenAIAdvancedSchedulerWeightTTFT + updates[SettingKeyOpenAIAdvancedSchedulerWeightReset] = settings.OpenAIAdvancedSchedulerWeightReset + updates[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom] = settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom + updates[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse] = settings.OpenAIAdvancedSchedulerWeightPreviousResponse + updates[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky] = settings.OpenAIAdvancedSchedulerWeightSessionSticky // 余额、订阅到期与账号限额通知 updates[SettingKeyBalanceLowNotifyEnabled] = strconv.FormatBool(settings.BalanceLowNotifyEnabled) @@ -2376,7 +2391,21 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { }) openAIAdvancedSchedulerSettingSF.Forget(openAIAdvancedSchedulerSettingKey) openAIAdvancedSchedulerSettingCache.Store(&cachedOpenAIAdvancedSchedulerSetting{ - enabled: settings.OpenAIAdvancedSchedulerEnabled, + enabled: settings.OpenAIAdvancedSchedulerEnabled, + stickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled, + subscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled, + lbTopKOverride: parsePositiveIntOverride(settings.OpenAIAdvancedSchedulerLBTopK), + weightOverrides: parseOpenAIAdvancedSchedulerWeightOverrides(map[string]string{ + SettingKeyOpenAIAdvancedSchedulerWeightPriority: settings.OpenAIAdvancedSchedulerWeightPriority, + SettingKeyOpenAIAdvancedSchedulerWeightLoad: settings.OpenAIAdvancedSchedulerWeightLoad, + SettingKeyOpenAIAdvancedSchedulerWeightQueue: settings.OpenAIAdvancedSchedulerWeightQueue, + SettingKeyOpenAIAdvancedSchedulerWeightErrorRate: settings.OpenAIAdvancedSchedulerWeightErrorRate, + SettingKeyOpenAIAdvancedSchedulerWeightTTFT: settings.OpenAIAdvancedSchedulerWeightTTFT, + SettingKeyOpenAIAdvancedSchedulerWeightReset: settings.OpenAIAdvancedSchedulerWeightReset, + SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, + SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: settings.OpenAIAdvancedSchedulerWeightPreviousResponse, + SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky: settings.OpenAIAdvancedSchedulerWeightSessionSticky, + }), expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(), }) // Invalidate the quota auto-pause cache and let the next read trigger a fresh load. @@ -3196,17 +3225,29 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { SettingKeyCodexCLIOnlyEngineFingerprintSignals: openai.DefaultEngineFingerprintSignalsJSON(), // 分组隔离(默认不允许未分组 Key 调度) - SettingKeyAllowUngroupedKeyScheduling: "false", - SettingKeyEnableAnthropicCacheTTL1hInjection: "false", - SettingKeyRewriteMessageCacheControl: strconv.FormatBool(s.defaultRewriteMessageCacheControl()), - SettingKeyEnableClientDatelineNormalization: "true", - SettingKeyAntigravityUserAgentVersion: "", - SettingKeyOpenAICodexUserAgent: "", - SettingPaymentVisibleMethodAlipaySource: "", - SettingPaymentVisibleMethodWxpaySource: "", - SettingPaymentVisibleMethodAlipayEnabled: "false", - SettingPaymentVisibleMethodWxpayEnabled: "false", - openAIAdvancedSchedulerSettingKey: "false", + SettingKeyAllowUngroupedKeyScheduling: "false", + SettingKeyEnableAnthropicCacheTTL1hInjection: "false", + SettingKeyRewriteMessageCacheControl: strconv.FormatBool(s.defaultRewriteMessageCacheControl()), + SettingKeyEnableClientDatelineNormalization: "true", + SettingKeyAntigravityUserAgentVersion: "", + SettingKeyOpenAICodexUserAgent: "", + SettingPaymentVisibleMethodAlipaySource: "", + SettingPaymentVisibleMethodWxpaySource: "", + SettingPaymentVisibleMethodAlipayEnabled: "false", + SettingPaymentVisibleMethodWxpayEnabled: "false", + openAIAdvancedSchedulerSettingKey: "false", + SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled: "false", + SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled: "false", + SettingKeyOpenAIAdvancedSchedulerLBTopK: "", + SettingKeyOpenAIAdvancedSchedulerWeightPriority: "", + SettingKeyOpenAIAdvancedSchedulerWeightLoad: "", + SettingKeyOpenAIAdvancedSchedulerWeightQueue: "", + SettingKeyOpenAIAdvancedSchedulerWeightErrorRate: "", + SettingKeyOpenAIAdvancedSchedulerWeightTTFT: "", + SettingKeyOpenAIAdvancedSchedulerWeightReset: "", + SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom: "", + SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: "", + SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky: "", SettingKeyAllowUserViewErrorRequests: "false", } @@ -3769,6 +3810,29 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin result.PaymentVisibleMethodAlipayEnabled = settings[SettingPaymentVisibleMethodAlipayEnabled] == "true" result.PaymentVisibleMethodWxpayEnabled = settings[SettingPaymentVisibleMethodWxpayEnabled] == "true" result.OpenAIAdvancedSchedulerEnabled = settings[openAIAdvancedSchedulerSettingKey] == "true" + result.OpenAIAdvancedSchedulerStickyWeightedEnabled = settings[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] == "true" + result.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled = settings[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] == "true" + result.OpenAIAdvancedSchedulerLBTopK = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + result.OpenAIAdvancedSchedulerWeightPriority = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightPriority]) + result.OpenAIAdvancedSchedulerWeightLoad = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightLoad]) + result.OpenAIAdvancedSchedulerWeightQueue = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightQueue]) + result.OpenAIAdvancedSchedulerWeightErrorRate = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightErrorRate]) + result.OpenAIAdvancedSchedulerWeightTTFT = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightTTFT]) + result.OpenAIAdvancedSchedulerWeightReset = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightReset]) + result.OpenAIAdvancedSchedulerWeightQuotaHeadroom = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom]) + result.OpenAIAdvancedSchedulerWeightPreviousResponse = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse]) + result.OpenAIAdvancedSchedulerWeightSessionSticky = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky]) + result.OpenAIAdvancedSchedulerEffectiveLBTopK = s.openAIAdvancedSchedulerEffectiveLBTopK() + effectiveWeights := s.openAIAdvancedSchedulerEffectiveWeights() + result.OpenAIAdvancedSchedulerEffectiveWeightPriority = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.Priority) + result.OpenAIAdvancedSchedulerEffectiveWeightLoad = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.Load) + result.OpenAIAdvancedSchedulerEffectiveWeightQueue = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.Queue) + result.OpenAIAdvancedSchedulerEffectiveWeightErrorRate = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.ErrorRate) + result.OpenAIAdvancedSchedulerEffectiveWeightTTFT = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.TTFT) + result.OpenAIAdvancedSchedulerEffectiveWeightReset = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.Reset) + result.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.QuotaHeadroom) + result.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.PreviousResponse) + result.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.SessionSticky) // 余额、订阅到期与账号限额通知 result.BalanceLowNotifyEnabled = settings[SettingKeyBalanceLowNotifyEnabled] == "true" @@ -3841,6 +3905,119 @@ func normalizeVisibleMethodSettingSource(method, source string, enabled bool) (s return normalized, nil } +func (s *SettingService) openAIAdvancedSchedulerEffectiveLBTopK() string { + if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.LBTopK > 0 { + return strconv.Itoa(s.cfg.Gateway.OpenAIWS.LBTopK) + } + return "7" +} + +func (s *SettingService) openAIAdvancedSchedulerEffectiveWeights() config.GatewayOpenAIWSSchedulerScoreWeights { + defaults := config.GatewayOpenAIWSSchedulerScoreWeights{ + Priority: 1.0, + Load: 1.0, + Queue: 0.7, + ErrorRate: 0.8, + TTFT: 0.5, + Reset: 0.0, + QuotaHeadroom: 0.0, + PreviousResponse: 5.0, + SessionSticky: 3.0, + } + if s == nil || s.cfg == nil { + return defaults + } + + weights := s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights + baseSum := weights.Priority + weights.Load + weights.Queue + weights.ErrorRate + weights.TTFT + weights.QuotaHeadroom + if baseSum <= 0 { + return defaults + } + return weights +} + +func formatOpenAIAdvancedSchedulerFloat(value float64) string { + return strconv.FormatFloat(value, 'f', -1, 64) +} + +func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { + lbTopK, err := normalizeOptionalPositiveIntString(settings.OpenAIAdvancedSchedulerLBTopK) + if err != nil { + return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_LB_TOP_K", "openai advanced scheduler TopK must be a positive integer or empty") + } + settings.OpenAIAdvancedSchedulerLBTopK = lbTopK + + weights := []*string{ + &settings.OpenAIAdvancedSchedulerWeightPriority, + &settings.OpenAIAdvancedSchedulerWeightLoad, + &settings.OpenAIAdvancedSchedulerWeightQueue, + &settings.OpenAIAdvancedSchedulerWeightErrorRate, + &settings.OpenAIAdvancedSchedulerWeightTTFT, + &settings.OpenAIAdvancedSchedulerWeightReset, + &settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, + &settings.OpenAIAdvancedSchedulerWeightPreviousResponse, + &settings.OpenAIAdvancedSchedulerWeightSessionSticky, + } + for _, target := range weights { + normalized, err := normalizeOptionalNonNegativeFloatString(*target) + if err != nil { + return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler weights must be non-negative numbers or empty") + } + *target = normalized + } + + // 与 config.Validate 的 "scheduler_score_weights must not all be zero" 保持一致: + // 覆盖值(空则回退到生效的配置值)叠加后的基础权重和不允许为 0, + // 否则调度会静默退化为 TopK 内均匀随机。 + effective := s.openAIAdvancedSchedulerEffectiveWeights() + baseSum := resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPriority, effective.Priority) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightLoad, effective.Load) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQueue, effective.Queue) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightErrorRate, effective.ErrorRate) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightTTFT, effective.TTFT) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, effective.QuotaHeadroom) + if baseSum <= 0 { + return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler base weights must not all be zero") + } + return nil +} + +// resolveOpenAIAdvancedSchedulerWeight 返回覆盖值(已归一化的非空字符串),空则回退默认值。 +func resolveOpenAIAdvancedSchedulerWeight(normalized string, fallback float64) float64 { + if normalized == "" { + return fallback + } + value, err := strconv.ParseFloat(normalized, 64) + if err != nil { + return fallback + } + return value +} + +func normalizeOptionalPositiveIntString(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", nil + } + value, err := strconv.Atoi(raw) + if err != nil || value <= 0 { + return "", fmt.Errorf("invalid positive integer") + } + return strconv.Itoa(value), nil +} + +func normalizeOptionalNonNegativeFloatString(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", nil + } + value, err := strconv.ParseFloat(raw, 64) + if err != nil || value < 0 || math.IsNaN(value) || math.IsInf(value, 0) { + return "", fmt.Errorf("invalid non-negative float") + } + return strconv.FormatFloat(value, 'f', -1, 64), nil +} + func parseDefaultSubscriptions(raw string) []DefaultSubscriptionSetting { raw = strings.TrimSpace(raw) if raw == "" { diff --git a/backend/internal/service/setting_service_update_test.go b/backend/internal/service/setting_service_update_test.go index 379bf9bc04..a47e1379f7 100644 --- a/backend/internal/service/setting_service_update_test.go +++ b/backend/internal/service/setting_service_update_test.go @@ -49,6 +49,42 @@ func (s *settingUpdateRepoStub) Delete(ctx context.Context, key string) error { panic("unexpected Delete call") } +type settingGetAllRepoStub struct { + values map[string]string +} + +func (s *settingGetAllRepoStub) Get(ctx context.Context, key string) (*Setting, error) { + panic("unexpected Get call") +} + +func (s *settingGetAllRepoStub) GetValue(ctx context.Context, key string) (string, error) { + panic("unexpected GetValue call") +} + +func (s *settingGetAllRepoStub) Set(ctx context.Context, key, value string) error { + panic("unexpected Set call") +} + +func (s *settingGetAllRepoStub) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) { + panic("unexpected GetMultiple call") +} + +func (s *settingGetAllRepoStub) SetMultiple(ctx context.Context, settings map[string]string) error { + panic("unexpected SetMultiple call") +} + +func (s *settingGetAllRepoStub) GetAll(ctx context.Context) (map[string]string, error) { + out := make(map[string]string, len(s.values)) + for key, value := range s.values { + out[key] = value + } + return out, nil +} + +func (s *settingGetAllRepoStub) Delete(ctx context.Context, key string) error { + panic("unexpected Delete call") +} + type settingAntigravityUARepoStub struct { values map[string]string } @@ -261,15 +297,30 @@ func TestSettingService_UpdateSettings_TablePreferences(t *testing.T) { } func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + defer resetOpenAIAdvancedSchedulerSettingCacheForTest() + repo := &settingUpdateRepoStub{} svc := NewSettingService(repo, &config.Config{}) err := svc.UpdateSettings(context.Background(), &SystemSettings{ - PaymentVisibleMethodAlipaySource: "alipay", - PaymentVisibleMethodWxpaySource: "easypay", - PaymentVisibleMethodAlipayEnabled: true, - PaymentVisibleMethodWxpayEnabled: false, - OpenAIAdvancedSchedulerEnabled: true, + PaymentVisibleMethodAlipaySource: "alipay", + PaymentVisibleMethodWxpaySource: "easypay", + PaymentVisibleMethodAlipayEnabled: true, + PaymentVisibleMethodWxpayEnabled: false, + OpenAIAdvancedSchedulerEnabled: true, + OpenAIAdvancedSchedulerStickyWeightedEnabled: true, + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: true, + OpenAIAdvancedSchedulerLBTopK: " 3 ", + OpenAIAdvancedSchedulerWeightPriority: "2.50", + OpenAIAdvancedSchedulerWeightLoad: "0", + OpenAIAdvancedSchedulerWeightQueue: "0.75", + OpenAIAdvancedSchedulerWeightErrorRate: "1.25", + OpenAIAdvancedSchedulerWeightTTFT: "0.5", + OpenAIAdvancedSchedulerWeightReset: "", + OpenAIAdvancedSchedulerWeightQuotaHeadroom: "0.2", + OpenAIAdvancedSchedulerWeightPreviousResponse: "8", + OpenAIAdvancedSchedulerWeightSessionSticky: "4", }) require.NoError(t, err) require.Equal(t, VisibleMethodSourceOfficialAlipay, repo.updates[SettingPaymentVisibleMethodAlipaySource]) @@ -277,6 +328,49 @@ func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler require.Equal(t, "true", repo.updates[SettingPaymentVisibleMethodAlipayEnabled]) require.Equal(t, "false", repo.updates[SettingPaymentVisibleMethodWxpayEnabled]) require.Equal(t, "true", repo.updates[openAIAdvancedSchedulerSettingKey]) + require.Equal(t, "true", repo.updates[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]) + require.Equal(t, "true", repo.updates[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]) + require.Equal(t, "3", repo.updates[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + require.Equal(t, "2.5", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightPriority]) + require.Equal(t, "0", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightLoad]) + require.Equal(t, "0.75", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightQueue]) + require.Equal(t, "1.25", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightErrorRate]) + require.Equal(t, "0.5", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightTTFT]) + require.Equal(t, "", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightReset]) + require.Equal(t, "0.2", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom]) + require.Equal(t, "8", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse]) + require.Equal(t, "4", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky]) +} + +func TestSettingService_GetAllSettings_OpenAIAdvancedSchedulerEffectiveValuesUseConfig(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 13 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights = config.GatewayOpenAIWSSchedulerScoreWeights{ + Priority: 2, + Load: 3, + Queue: 4, + ErrorRate: 5, + TTFT: 6, + Reset: 7, + QuotaHeadroom: 8, + PreviousResponse: 9, + SessionSticky: 10, + } + svc := NewSettingService(&settingGetAllRepoStub{values: map[string]string{ + SettingKeyOpenAIAdvancedSchedulerLBTopK: "3", + SettingKeyOpenAIAdvancedSchedulerWeightPriority: "99", + SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky: "88", + }}, cfg) + + settings, err := svc.GetAllSettings(context.Background()) + require.NoError(t, err) + require.Equal(t, "3", settings.OpenAIAdvancedSchedulerLBTopK) + require.Equal(t, "99", settings.OpenAIAdvancedSchedulerWeightPriority) + require.Equal(t, "88", settings.OpenAIAdvancedSchedulerWeightSessionSticky) + require.Equal(t, "13", settings.OpenAIAdvancedSchedulerEffectiveLBTopK) + require.Equal(t, "2", settings.OpenAIAdvancedSchedulerEffectiveWeightPriority) + require.Equal(t, "3", settings.OpenAIAdvancedSchedulerEffectiveWeightLoad) + require.Equal(t, "10", settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky) } func TestSettingService_UpdateSettings_AntigravityUserAgentVersion(t *testing.T) { diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index ac225e1b14..465760a445 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -218,7 +218,29 @@ type SystemSettings struct { PaymentVisibleMethodWxpayEnabled bool // OpenAI 账号调度 - OpenAIAdvancedSchedulerEnabled bool + OpenAIAdvancedSchedulerEnabled bool + OpenAIAdvancedSchedulerStickyWeightedEnabled bool + OpenAIAdvancedSchedulerSubscriptionPriorityEnabled bool + OpenAIAdvancedSchedulerLBTopK string + OpenAIAdvancedSchedulerWeightPriority string + OpenAIAdvancedSchedulerWeightLoad string + OpenAIAdvancedSchedulerWeightQueue string + OpenAIAdvancedSchedulerWeightErrorRate string + OpenAIAdvancedSchedulerWeightTTFT string + OpenAIAdvancedSchedulerWeightReset string + OpenAIAdvancedSchedulerWeightQuotaHeadroom string + OpenAIAdvancedSchedulerWeightPreviousResponse string + OpenAIAdvancedSchedulerWeightSessionSticky string + OpenAIAdvancedSchedulerEffectiveLBTopK string + OpenAIAdvancedSchedulerEffectiveWeightPriority string + OpenAIAdvancedSchedulerEffectiveWeightLoad string + OpenAIAdvancedSchedulerEffectiveWeightQueue string + OpenAIAdvancedSchedulerEffectiveWeightErrorRate string + OpenAIAdvancedSchedulerEffectiveWeightTTFT string + OpenAIAdvancedSchedulerEffectiveWeightReset string + OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom string + OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse string + OpenAIAdvancedSchedulerEffectiveWeightSessionSticky string // 余额不足提醒 BalanceLowNotifyEnabled bool diff --git a/backend/internal/service/token_refresh_service.go b/backend/internal/service/token_refresh_service.go index 08761f8220..2d4a026692 100644 --- a/backend/internal/service/token_refresh_service.go +++ b/backend/internal/service/token_refresh_service.go @@ -312,6 +312,7 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc if isNonRetryableRefreshError(err) { errorMsg := "Token refresh failed (non-retryable): " + logredact.RedactText(err.Error()) s.notifyAccountSchedulingBlocked(account, time.Time{}, "token_refresh_non_retryable") + s.clearAntigravityForceTokenRefresh(ctx, account, "non_retryable") if setErr := s.accountRepo.SetError(ctx, account.ID, errorMsg); setErr != nil { slog.Error("token_refresh.set_error_status_failed", "account_id", account.ID, @@ -369,6 +370,8 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc // postRefreshActions 刷新成功后的后续动作(清除错误状态、缓存失效、调度器同步等) func (s *TokenRefreshService) postRefreshActions(ctx context.Context, account *Account) { + s.clearAntigravityForceTokenRefresh(ctx, account, "success") + // Antigravity 账户:如果之前是因为缺少 project_id 而标记为 error,现在成功获取到了,清除错误状态 if account.Platform == PlatformAntigravity && account.Status == StatusError && @@ -432,6 +435,30 @@ func (s *TokenRefreshService) postRefreshActions(ctx context.Context, account *A s.ensureAntigravityPrivacy(ctx, account) } +func (s *TokenRefreshService) clearAntigravityForceTokenRefresh(ctx context.Context, account *Account, outcome string) { + if s == nil || account == nil || !accountNeedsAntigravityForceTokenRefresh(account) { + return + } + updates := clearAntigravityForceTokenRefreshExtra() + if err := s.accountRepo.UpdateExtra(ctx, account.ID, updates); err != nil { + slog.Warn("token_refresh.clear_antigravity_force_refresh_failed", + "account_id", account.ID, + "outcome", outcome, + "error", err, + ) + return + } + if account.Extra != nil { + for k, v := range updates { + account.Extra[k] = v + } + } + slog.Info("token_refresh.cleared_antigravity_force_refresh", + "account_id", account.ID, + "outcome", outcome, + ) +} + // errRefreshSkipped 表示刷新被跳过(锁竞争或已被其他路径刷新),不计入 failed 或 refreshed var errRefreshSkipped = fmt.Errorf("refresh skipped") @@ -446,6 +473,7 @@ func isNonRetryableRefreshError(err error) bool { nonRetryable := []string{ "invalid_grant", // refresh_token 已失效 "invalid_refresh_token", // refresh_token 无效, team 账号工作区被删除会出现 + "token_expired", // OpenAI refresh_token 已过期,需要重新授权 "app_session_terminated", // refresh_token team 账号工作区被删除 "refresh_token_reused", // OpenAI refresh_token 已被使用,必须重新授权 "refresh_token_invalidated", // OpenAI session ended; refresh token invalidated diff --git a/backend/internal/service/token_refresh_service_test.go b/backend/internal/service/token_refresh_service_test.go index d2315c1db5..bd6a521576 100644 --- a/backend/internal/service/token_refresh_service_test.go +++ b/backend/internal/service/token_refresh_service_test.go @@ -20,8 +20,10 @@ type tokenRefreshAccountRepo struct { setErrorCalls int clearTempCalls int setTempUnschedCalls int + updateExtraCalls int lastErrorMessage string lastTempUnschedReason string + lastExtraUpdates map[string]any lastAccount *Account updateErr error } @@ -68,6 +70,22 @@ func (r *tokenRefreshAccountRepo) SetTempUnschedulable(ctx context.Context, id i return nil } +func (r *tokenRefreshAccountRepo) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error { + r.updateExtraCalls++ + r.lastExtraUpdates = shallowCopyMap(updates) + if r.accountsByID != nil { + if acc, ok := r.accountsByID[id]; ok && acc != nil { + if acc.Extra == nil { + acc.Extra = make(map[string]any, len(updates)) + } + for k, v := range updates { + acc.Extra[k] = v + } + } + } + return nil +} + type tokenCacheInvalidatorStub struct { calls int err error @@ -233,6 +251,121 @@ func TestTokenRefreshService_RefreshWithRetry_Antigravity(t *testing.T) { require.Equal(t, 1, invalidator.calls) // Antigravity 也应触发缓存失效 } +func TestAntigravityTokenRefresher_NeedsRefresh_ForceRefreshMarker(t *testing.T) { + refresher := NewAntigravityTokenRefresher(nil) + account := &Account{ + ID: 3675, + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "expires_at": time.Now().Add(time.Hour).Format(time.RFC3339), + }, + Extra: map[string]any{ + antigravityForceTokenRefreshExtraKey: true, + }, + } + + require.True(t, refresher.NeedsRefresh(account, 0), "server-invalidated token must refresh even before expires_at") +} + +func TestAntigravityTokenRefresher_NeedsRefresh_NormalExpiryRulesUnchanged(t *testing.T) { + refresher := NewAntigravityTokenRefresher(nil) + + t.Run("normal_unexpired_without_marker_does_not_refresh", func(t *testing.T) { + account := &Account{ + ID: 3707, + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "expires_at": time.Now().Add(time.Hour).Format(time.RFC3339), + }, + } + + require.False(t, refresher.NeedsRefresh(account, 0)) + }) + + t.Run("normal_expiring_refreshes", func(t *testing.T) { + account := &Account{ + ID: 3708, + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "expires_at": time.Now().Add(5 * time.Minute).Format(time.RFC3339), + }, + } + + require.True(t, refresher.NeedsRefresh(account, 0)) + }) +} + +func TestTokenRefreshService_RefreshWithRetry_AntigravityClearsForceRefreshOnSuccess(t *testing.T) { + repo := &tokenRefreshAccountRepo{} + cfg := &config.Config{ + TokenRefresh: config.TokenRefreshConfig{ + MaxRetries: 1, + RetryBackoffSeconds: 0, + }, + } + service := NewTokenRefreshService(repo, nil, nil, nil, nil, nil, nil, cfg, nil) + until := time.Now().Add(10 * time.Minute) + account := &Account{ + ID: 3709, + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + TempUnschedulableUntil: &until, + Extra: map[string]any{ + antigravityForceTokenRefreshExtraKey: true, + antigravityForceTokenRefreshReasonExtraKey: "401_invalid", + "privacy_mode": AntigravityPrivacySet, + }, + } + refresher := &tokenRefresherStub{ + credentials: map[string]any{ + "access_token": "new-ag-token", + }, + } + + err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour) + require.NoError(t, err) + require.Equal(t, 1, repo.updateCredentialsCalls) + require.Equal(t, 1, repo.updateExtraCalls) + require.Equal(t, false, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey]) + require.Equal(t, "", repo.lastExtraUpdates[antigravityForceTokenRefreshReasonExtraKey]) + require.Equal(t, false, account.Extra[antigravityForceTokenRefreshExtraKey]) + require.Equal(t, 1, repo.clearTempCalls, "successful refresh should restore schedulability") +} + +func TestTokenRefreshService_RefreshWithRetry_AntigravityForceRefreshInvalidGrantSetsError(t *testing.T) { + repo := &tokenRefreshAccountRepo{} + cfg := &config.Config{ + TokenRefresh: config.TokenRefreshConfig{ + MaxRetries: 3, + RetryBackoffSeconds: 0, + }, + } + service := NewTokenRefreshService(repo, nil, nil, nil, nil, nil, nil, cfg, nil) + account := &Account{ + ID: 3710, + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Extra: map[string]any{ + antigravityForceTokenRefreshExtraKey: true, + antigravityForceTokenRefreshReasonExtraKey: "401_invalid", + }, + } + refresher := &tokenRefresherStub{ + err: errors.New("invalid_grant: token revoked"), + } + + err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour) + require.Error(t, err) + require.Equal(t, 1, repo.setErrorCalls) + require.Equal(t, 0, repo.setTempUnschedCalls) + require.Equal(t, 1, repo.updateExtraCalls) + require.Equal(t, false, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey]) + require.Contains(t, repo.lastErrorMessage, "non-retryable") +} + // TestTokenRefreshService_RefreshWithRetry_NonOAuthAccount 测试非 OAuth 账号不触发缓存失效 func TestTokenRefreshService_RefreshWithRetry_NonOAuthAccount(t *testing.T) { repo := &tokenRefreshAccountRepo{} @@ -541,6 +674,7 @@ func TestIsNonRetryableRefreshError(t *testing.T) { {name: "invalid_grant", err: errors.New("invalid_grant"), expected: true}, {name: "invalid_client", err: errors.New("invalid_client"), expected: true}, {name: "invalid_refresh_token", err: errors.New(`OPENAI_OAUTH_TOKEN_REFRESH_FAILED: token refresh failed: status 401, body: {"error":{"code":"invalid_refresh_token"}}`), expected: true}, + {name: "token_expired", err: errors.New(`OPENAI_OAUTH_TOKEN_REFRESH_FAILED: token refresh failed: status 401, body: {"error":{"code":"token_expired"}}`), expected: true}, {name: "refresh_token_reused", err: errors.New(`OPENAI_OAUTH_TOKEN_REFRESH_FAILED: token refresh failed: status 401, body: {"error":{"code":"refresh_token_reused"}}`), expected: true}, {name: "app_session_terminated", err: errors.New(`OPENAI_OAUTH_TOKEN_REFRESH_FAILED: token refresh failed: status 401, body: {"error": {"code": "app_session_terminated"}}`), expected: true}, {name: "unauthorized_client", err: errors.New("unauthorized_client"), expected: true}, diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index ee3e6bfc04..4f6a305b25 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -208,6 +208,8 @@ func (s *AccountTestService) buildAnthropicUpstreamModelsRequest(ctx context.Con } else { setAnthropicAPIKeyAuthHeader(req.Header, account, apiKeyAuthToken) } + // 账号级请求头覆写:模型列表探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) return req, nil } @@ -277,6 +279,8 @@ func (s *AccountTestService) buildOpenAIUpstreamModelsRequest(ctx context.Contex } req.Header.Set("Accept", "application/json") req.Header.Set("Authorization", "Bearer "+apiKey) + // 账号级请求头覆写:模型列表探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) return req, nil } @@ -387,14 +391,7 @@ func buildV1ModelsURL(base string) string { } func buildOpenAIModelsURL(base string) string { - normalized := strings.TrimRight(strings.TrimSpace(base), "/") - if strings.HasSuffix(normalized, "/v1/models") { - return normalized - } - if strings.HasSuffix(normalized, "/v1") { - return normalized + "/models" - } - return normalized + "/v1/models" + return buildOpenAIEndpointURL(base, "/v1/models") } func buildGeminiModelsURL(base string) string { diff --git a/backend/internal/service/upstream_models_test.go b/backend/internal/service/upstream_models_test.go index 1fe9415d34..3904194ffa 100644 --- a/backend/internal/service/upstream_models_test.go +++ b/backend/internal/service/upstream_models_test.go @@ -29,6 +29,61 @@ func TestBuildV1ModelsURL(t *testing.T) { require.Equal(t, "https://gateway.example.com/antigravity/v1/models", buildV1ModelsURL("https://gateway.example.com/antigravity/")) } +func TestBuildOpenAIModelsURL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + base string + want string + }{ + { + name: "zhipu v4 coding base url", + base: "https://open.bigmodel.cn/api/coding/paas/v4", + want: "https://open.bigmodel.cn/api/coding/paas/v4/models", + }, + { + name: "openai v1 base url", + base: "https://api.openai.com/v1", + want: "https://api.openai.com/v1/models", + }, + { + name: "models url unchanged", + base: "https://api.openai.com/v1/models", + want: "https://api.openai.com/v1/models", + }, + { + name: "host fallback uses v1", + base: "https://api.openai.com", + want: "https://api.openai.com/v1/models", + }, + { + name: "trailing slash on v4", + base: "https://open.bigmodel.cn/api/coding/paas/v4/", + want: "https://open.bigmodel.cn/api/coding/paas/v4/models", + }, + { + name: "v2 base url", + base: "https://gateway.example.com/openai/v2", + want: "https://gateway.example.com/openai/v2/models", + }, + { + name: "v3 base url", + base: "https://gateway.example.com/openai/v3", + want: "https://gateway.example.com/openai/v3/models", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + require.Equal(t, tt.want, buildOpenAIModelsURL(tt.base)) + }) + } +} + func TestBuildGeminiModelsURL(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/usage_record_worker_pool.go b/backend/internal/service/usage_record_worker_pool.go index 5da0b89023..bb5ae452c8 100644 --- a/backend/internal/service/usage_record_worker_pool.go +++ b/backend/internal/service/usage_record_worker_pool.go @@ -15,10 +15,11 @@ import ( ) const ( - defaultUsageRecordWorkerCount = 128 - defaultUsageRecordQueueSize = 16384 - defaultUsageRecordTaskTimeoutSeconds = 5 - defaultUsageRecordOverflowPolicy = config.UsageRecordOverflowPolicySample + defaultUsageRecordWorkerCount = 128 + defaultUsageRecordQueueSize = 16384 + defaultUsageRecordTaskTimeoutSeconds = 5 + // 默认 sync:溢出时提交方内联执行,保证计费任务不被静默丢弃(issue #3656)。 + defaultUsageRecordOverflowPolicy = config.UsageRecordOverflowPolicySync defaultUsageRecordOverflowSampleRatio = 10 defaultUsageRecordAutoScaleEnabled = true defaultUsageRecordAutoScaleMinWorkers = 128 diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 7278a4e0c5..908783eb88 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -537,9 +537,11 @@ func ProvideAPIKeyService( cache APIKeyCache, cfg *config.Config, billingCacheService *BillingCacheService, + concurrencyService *ConcurrencyService, ) *APIKeyService { svc := NewAPIKeyService(apiKeyRepo, userRepo, groupRepo, userSubRepo, userGroupRateRepo, cache, cfg) svc.SetRateLimitCacheInvalidator(billingCacheService) + svc.SetConcurrencyService(concurrencyService) return svc } diff --git a/backend/internal/setup/setup.go b/backend/internal/setup/setup.go index 51baf3dfe6..a2c4e2847a 100644 --- a/backend/internal/setup/setup.go +++ b/backend/internal/setup/setup.go @@ -28,6 +28,7 @@ const ( InstallLockFile = ".installed" defaultUserConcurrency = 5 simpleModeAdminConcurrency = 30 + defaultMigrationTimeout = 60 * time.Second ) func setupDefaultAdminConcurrency() int { @@ -73,12 +74,13 @@ func GetInstallLockPath() string { // SetupConfig holds the setup configuration type SetupConfig struct { - Database DatabaseConfig `json:"database" yaml:"database"` - Redis RedisConfig `json:"redis" yaml:"redis"` - Admin AdminConfig `json:"admin" yaml:"-"` // Not stored in config file - Server ServerConfig `json:"server" yaml:"server"` - JWT JWTConfig `json:"jwt" yaml:"jwt"` - Timezone string `json:"timezone" yaml:"timezone"` // e.g. "Asia/Shanghai", "UTC" + Database DatabaseConfig `json:"database" yaml:"database"` + Redis RedisConfig `json:"redis" yaml:"redis"` + Admin AdminConfig `json:"admin" yaml:"-"` // Not stored in config file + Server ServerConfig `json:"server" yaml:"server"` + JWT JWTConfig `json:"jwt" yaml:"jwt"` + Timezone string `json:"timezone" yaml:"timezone"` // e.g. "Asia/Shanghai", "UTC" + MigrationTimeoutSeconds int `json:"migration_timeout_seconds" yaml:"migration_timeout_seconds,omitempty"` } type DatabaseConfig struct { @@ -350,11 +352,18 @@ func initializeDatabase(cfg *SetupConfig) error { } }() - migrationCtx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + migrationCtx, cancel := context.WithTimeout(context.Background(), cfg.migrationTimeout()) defer cancel() return repository.ApplyMigrations(migrationCtx, db) } +func (cfg *SetupConfig) migrationTimeout() time.Duration { + if cfg != nil && cfg.MigrationTimeoutSeconds > 0 { + return time.Duration(cfg.MigrationTimeoutSeconds) * time.Second + } + return defaultMigrationTimeout +} + func createAdminUser(cfg *SetupConfig) (bool, string, error) { dsn := fmt.Sprintf( "host=%s port=%d user=%s password=%s dbname=%s sslmode=%s", @@ -578,7 +587,8 @@ func AutoSetupFromEnv() error { Secret: getEnvOrDefault("JWT_SECRET", ""), ExpireHour: getEnvIntOrDefault("JWT_EXPIRE_HOUR", 24), }, - Timezone: tz, + Timezone: tz, + MigrationTimeoutSeconds: getEnvIntOrDefault("SETUP_MIGRATION_TIMEOUT_SECONDS", 0), } // Generate JWT secret if not provided diff --git a/backend/internal/setup/setup_test.go b/backend/internal/setup/setup_test.go index a2aa2f4cc1..b95c162bda 100644 --- a/backend/internal/setup/setup_test.go +++ b/backend/internal/setup/setup_test.go @@ -4,6 +4,7 @@ import ( "os" "strings" "testing" + "time" ) func TestDecideAdminBootstrap(t *testing.T) { @@ -70,6 +71,22 @@ func TestSetupDefaultAdminConcurrency(t *testing.T) { }) } +func TestSetupMigrationTimeout(t *testing.T) { + t.Run("uses default timeout when unset", func(t *testing.T) { + cfg := &SetupConfig{} + if got := cfg.migrationTimeout(); got != 60*time.Second { + t.Fatalf("migrationTimeout()=%s, want 60s", got) + } + }) + + t.Run("uses configured timeout", func(t *testing.T) { + cfg := &SetupConfig{MigrationTimeoutSeconds: 300} + if got := cfg.migrationTimeout(); got != 300*time.Second { + t.Fatalf("migrationTimeout()=%s, want 300s", got) + } + }) +} + func TestWriteConfigFileKeepsDefaultUserConcurrency(t *testing.T) { t.Setenv("RUN_MODE", "simple") t.Setenv("DATA_DIR", t.TempDir()) diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index e88ed2da22..e7d4e7ded3 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -4886,6 +4886,150 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "gpt-5.6-sol": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.6-terra": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.6-luna": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "gpt-5.5": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_272k_tokens": 1e-06, diff --git a/deploy/.env.example b/deploy/.env.example index 59e4b44b91..5925f0abb4 100644 --- a/deploy/.env.example +++ b/deploy/.env.example @@ -196,6 +196,13 @@ JWT_EXPIRE_HOUR=24 # - =0: 回退使用 JWT_EXPIRE_HOUR JWT_ACCESS_TOKEN_EXPIRE_MINUTES=0 +# ----------------------------------------------------------------------------- +# Setup Configuration +# ----------------------------------------------------------------------------- +# Database migration timeout during initial setup, in seconds. +# Leave 0 to use the built-in default of 60 seconds. +SETUP_MIGRATION_TIMEOUT_SECONDS=0 + # ----------------------------------------------------------------------------- # TOTP (2FA) Configuration # TOTP(双因素认证)配置 @@ -352,11 +359,11 @@ DASHBOARD_AGGREGATION_RETENTION_DAILY_DAYS=730 # 启用 URL 白名单验证(false 则跳过白名单检查,仅做基本格式校验) SECURITY_URL_ALLOWLIST_ENABLED=false -# 关闭白名单时,是否允许 http:// URL(默认 false,只允许 https://) -# ⚠️ 警告:允许 HTTP 存在安全风险(明文传输),仅建议在开发/测试环境或可信内网中使用 -# Allow insecure HTTP URLs when allowlist is disabled (default: false, requires https) +# 关闭白名单时,是否允许 http:// URL(默认 true,设为 false 则只允许 https://) +# ⚠️ 警告:允许 HTTP 存在安全风险(明文传输),生产环境建议设为 false +# Allow insecure HTTP URLs when allowlist is disabled (default: true; set to false to require https) # ⚠️ WARNING: Allowing HTTP has security risks (plaintext transmission) -# Only recommended for dev/test environments or trusted networks +# Recommended to set false in production SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=true # 是否允许本地/私有 IP 地址用于上游/定价/CRS(仅在可信网络中使用) diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index 5a3afb0315..eff4bfb598 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -108,8 +108,8 @@ security: # Allow localhost/private IPs for upstream/pricing/CRS (use only in trusted networks) # 允许本地/私有 IP 地址用于上游/定价/CRS(仅在可信网络中使用) allow_private_hosts: true - # Allow http:// URLs when allowlist is disabled (default: false, require https) - # 白名单禁用时是否允许 http:// URL(默认: false,要求 https) + # Allow http:// URLs when allowlist is disabled (default: true; set to false to require https) + # 白名单禁用时是否允许 http:// URL(默认: true,设为 false 则仅允许 https) allow_insecure_http: true response_headers: # Enable configurable response header filtering (default: true) diff --git a/deploy/docker-compose.dev.yml b/deploy/docker-compose.dev.yml index 7755fdbeff..07e89e0b51 100644 --- a/deploy/docker-compose.dev.yml +++ b/deploy/docker-compose.dev.yml @@ -38,6 +38,7 @@ services: - ADMIN_EMAIL=${ADMIN_EMAIL:-admin@sub2api.local} - ADMIN_PASSWORD=${ADMIN_PASSWORD:-} - JWT_SECRET=${JWT_SECRET:-} + - SETUP_MIGRATION_TIMEOUT_SECONDS=${SETUP_MIGRATION_TIMEOUT_SECONDS:-0} - TOTP_ENCRYPTION_KEY=${TOTP_ENCRYPTION_KEY:-} - TZ=${TZ:-Asia/Shanghai} # OpenAI HTTP upstream protocol/timeout diff --git a/deploy/docker-compose.local.yml b/deploy/docker-compose.local.yml index b15be2402d..042752e857 100644 --- a/deploy/docker-compose.local.yml +++ b/deploy/docker-compose.local.yml @@ -94,6 +94,11 @@ services: - JWT_SECRET=${JWT_SECRET:-} - JWT_EXPIRE_HOUR=${JWT_EXPIRE_HOUR:-24} + # ======================================================================= + # Setup Configuration + # ======================================================================= + - SETUP_MIGRATION_TIMEOUT_SECONDS=${SETUP_MIGRATION_TIMEOUT_SECONDS:-0} + # ======================================================================= # TOTP (2FA) Configuration # ======================================================================= @@ -134,10 +139,10 @@ services: # ======================================================================= # Enable URL allowlist validation (false to skip allowlist checks) - SECURITY_URL_ALLOWLIST_ENABLED=${SECURITY_URL_ALLOWLIST_ENABLED:-false} - # Allow insecure HTTP URLs when allowlist is disabled (default: false, requires https) - - SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=${SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP:-false} - # Allow private IP addresses for upstream/pricing/CRS (for internal deployments) - - SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS=${SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS:-false} + # Allow insecure HTTP URLs when allowlist is disabled (default: true; set to false to require https) + - SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=${SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP:-true} + # Allow private IP addresses for upstream/pricing/CRS (default: true; set to false to block private hosts) + - SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS=${SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS:-true} # Upstream hosts whitelist (comma-separated, only used when enabled=true) - SECURITY_URL_ALLOWLIST_UPSTREAM_HOSTS=${SECURITY_URL_ALLOWLIST_UPSTREAM_HOSTS:-} diff --git a/deploy/docker-compose.standalone.yml b/deploy/docker-compose.standalone.yml index 32afb28d6c..2e1d335624 100644 --- a/deploy/docker-compose.standalone.yml +++ b/deploy/docker-compose.standalone.yml @@ -76,6 +76,11 @@ services: - JWT_SECRET=${JWT_SECRET:-} - JWT_EXPIRE_HOUR=${JWT_EXPIRE_HOUR:-24} + # ======================================================================= + # Setup Configuration + # ======================================================================= + - SETUP_MIGRATION_TIMEOUT_SECONDS=${SETUP_MIGRATION_TIMEOUT_SECONDS:-0} + # ======================================================================= # Timezone Configuration # ======================================================================= diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index fd682a87e6..22713c59aa 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -90,6 +90,11 @@ services: - JWT_SECRET=${JWT_SECRET:-} - JWT_EXPIRE_HOUR=${JWT_EXPIRE_HOUR:-24} + # ======================================================================= + # Setup Configuration + # ======================================================================= + - SETUP_MIGRATION_TIMEOUT_SECONDS=${SETUP_MIGRATION_TIMEOUT_SECONDS:-0} + # ======================================================================= # TOTP (2FA) Configuration # ======================================================================= @@ -130,10 +135,10 @@ services: # ======================================================================= # Enable URL allowlist validation (false to skip allowlist checks) - SECURITY_URL_ALLOWLIST_ENABLED=${SECURITY_URL_ALLOWLIST_ENABLED:-false} - # Allow insecure HTTP URLs when allowlist is disabled (default: false, requires https) - - SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=${SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP:-false} - # Allow private IP addresses for upstream/pricing/CRS (for internal deployments) - - SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS=${SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS:-false} + # Allow insecure HTTP URLs when allowlist is disabled (default: true; set to false to require https) + - SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP=${SECURITY_URL_ALLOWLIST_ALLOW_INSECURE_HTTP:-true} + # Allow private IP addresses for upstream/pricing/CRS (default: true; set to false to block private hosts) + - SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS=${SECURITY_URL_ALLOWLIST_ALLOW_PRIVATE_HOSTS:-true} # Upstream hosts whitelist (comma-separated, only used when enabled=true) - SECURITY_URL_ALLOWLIST_UPSTREAM_HOSTS=${SECURITY_URL_ALLOWLIST_UPSTREAM_HOSTS:-} diff --git a/frontend/src/__tests__/setup.ts b/frontend/src/__tests__/setup.ts index b777b22e8e..9dad8c1f12 100644 --- a/frontend/src/__tests__/setup.ts +++ b/frontend/src/__tests__/setup.ts @@ -57,6 +57,20 @@ if (typeof globalThis.cancelIdleCallback === 'undefined') { }) as unknown as typeof cancelIdleCallback } +// Mock matchMedia (jsdom 未实现;DataTable 等组件依赖它做桌面/移动分支) +if (typeof window !== 'undefined' && typeof window.matchMedia !== 'function') { + window.matchMedia = ((query: string) => ({ + matches: true, // 测试默认按桌面视口渲染表格 + media: query, + onchange: null, + addListener: vi.fn(), + removeListener: vi.fn(), + addEventListener: vi.fn(), + removeEventListener: vi.fn(), + dispatchEvent: vi.fn(), + })) as unknown as typeof window.matchMedia +} + // Mock IntersectionObserver class MockIntersectionObserver { observe = vi.fn() diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts index b3e53893ce..c7cbc64a4b 100644 --- a/frontend/src/api/admin/ops.ts +++ b/frontend/src/api/admin/ops.ts @@ -930,11 +930,16 @@ export interface OpsErrorLog { requested_model?: string upstream_model?: string request_type?: number | null + user_agent?: string + + // 已删除 KEY 所有者(INVALID_API_KEY 归因快照):认证失败行 user_id 为空, + // 用户列以此回退显示所有者 + deleted_key_owner_user_id?: number | null + deleted_key_owner_email?: string | null } export interface OpsErrorDetail extends OpsErrorLog { error_body: string - user_agent: string // Upstream context (optional; enriched by gateway services) upstream_status_code?: number | null @@ -950,10 +955,9 @@ export interface OpsErrorDetail extends OpsErrorLog { is_business_limited: boolean - // Deleted key owner info (INVALID_API_KEY attribution) + // Deleted key owner info (INVALID_API_KEY attribution); + // owner user_id/email 已上移到 OpsErrorLog(列表用户列回退) attempted_key_prefix?: string | null - deleted_key_owner_user_id?: number | null - deleted_key_owner_email?: string | null deleted_key_name?: string | null // Bound (non-deleted) key prefix, snapshotted at error time @@ -1098,6 +1102,8 @@ export type OpsErrorListQueryParams = { model?: string phase?: string + // 分类(用户侧粗分类码,如 auth/rate_limit/upstream),后端反查为 phase/type ANY 条件 + category?: string error_owner?: string error_source?: string resolved?: string @@ -1106,6 +1112,10 @@ export type OpsErrorListQueryParams = { q?: string status_codes?: string status_codes_other?: string + + // 服务端排序,列白名单见后端 opsErrorLogsOrderBy(created_at/model/status_code) + sort_by?: string + sort_order?: 'asc' | 'desc' } // Legacy unified endpoints diff --git a/frontend/src/api/admin/payment.ts b/frontend/src/api/admin/payment.ts index 49efcc355d..1d4305948e 100644 --- a/frontend/src/api/admin/payment.ts +++ b/frontend/src/api/admin/payment.ts @@ -24,6 +24,8 @@ export interface AdminPaymentConfig { enabled_payment_types: string[] balance_disabled: boolean balance_recharge_multiplier: number + subscription_usd_to_cny_rate: number + recharge_fee_rate: number load_balance_strategy: string product_name_prefix: string product_name_suffix: string @@ -42,6 +44,8 @@ export interface UpdatePaymentConfigRequest { enabled_payment_types?: string[] balance_disabled?: boolean balance_recharge_multiplier?: number + subscription_usd_to_cny_rate?: number + recharge_fee_rate?: number load_balance_strategy?: string product_name_prefix?: string product_name_suffix?: string diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 44fbe29187..f5da990930 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -589,6 +589,7 @@ export interface SystemSettings { payment_enabled_types: string[]; payment_balance_disabled: boolean; payment_balance_recharge_multiplier: number; + payment_subscription_usd_to_cny_rate: number; payment_recharge_fee_rate: number; payment_load_balance_strategy: string; payment_product_name_prefix: string; @@ -606,6 +607,28 @@ export interface SystemSettings { payment_visible_method_alipay_enabled?: boolean; payment_visible_method_wxpay_enabled?: boolean; openai_advanced_scheduler_enabled?: boolean; + openai_advanced_scheduler_sticky_weighted_enabled?: boolean; + openai_advanced_scheduler_subscription_priority_enabled?: boolean; + openai_advanced_scheduler_lb_top_k?: string; + openai_advanced_scheduler_weight_priority?: string; + openai_advanced_scheduler_weight_load?: string; + openai_advanced_scheduler_weight_queue?: string; + openai_advanced_scheduler_weight_error_rate?: string; + openai_advanced_scheduler_weight_ttft?: string; + openai_advanced_scheduler_weight_reset?: string; + openai_advanced_scheduler_weight_quota_headroom?: string; + openai_advanced_scheduler_weight_previous_response?: string; + openai_advanced_scheduler_weight_session_sticky?: string; + openai_advanced_scheduler_effective_lb_top_k?: string; + openai_advanced_scheduler_effective_weight_priority?: string; + openai_advanced_scheduler_effective_weight_load?: string; + openai_advanced_scheduler_effective_weight_queue?: string; + openai_advanced_scheduler_effective_weight_error_rate?: string; + openai_advanced_scheduler_effective_weight_ttft?: string; + openai_advanced_scheduler_effective_weight_reset?: string; + openai_advanced_scheduler_effective_weight_quota_headroom?: string; + openai_advanced_scheduler_effective_weight_previous_response?: string; + openai_advanced_scheduler_effective_weight_session_sticky?: string; // 余额、订阅到期与账号限额通知 balance_low_notify_enabled: boolean; @@ -838,6 +861,7 @@ export interface UpdateSettingsRequest { payment_enabled_types?: string[]; payment_balance_disabled?: boolean; payment_balance_recharge_multiplier?: number; + payment_subscription_usd_to_cny_rate?: number; payment_recharge_fee_rate?: number; payment_load_balance_strategy?: string; payment_product_name_prefix?: string; @@ -855,6 +879,18 @@ export interface UpdateSettingsRequest { payment_visible_method_alipay_enabled?: boolean; payment_visible_method_wxpay_enabled?: boolean; openai_advanced_scheduler_enabled?: boolean; + openai_advanced_scheduler_sticky_weighted_enabled?: boolean; + openai_advanced_scheduler_subscription_priority_enabled?: boolean; + openai_advanced_scheduler_lb_top_k?: string; + openai_advanced_scheduler_weight_priority?: string; + openai_advanced_scheduler_weight_load?: string; + openai_advanced_scheduler_weight_queue?: string; + openai_advanced_scheduler_weight_error_rate?: string; + openai_advanced_scheduler_weight_ttft?: string; + openai_advanced_scheduler_weight_reset?: string; + openai_advanced_scheduler_weight_quota_headroom?: string; + openai_advanced_scheduler_weight_previous_response?: string; + openai_advanced_scheduler_weight_session_sticky?: string; // 余额、订阅到期与账号限额通知 balance_low_notify_enabled?: boolean; balance_low_notify_threshold?: number; diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts index b37996d61a..83be033c11 100644 --- a/frontend/src/api/admin/usage.ts +++ b/frontend/src/api/admin/usage.ts @@ -86,6 +86,10 @@ export interface AdminUsageQueryParams extends UsageQueryParams { billing_mode?: string sort_by?: string sort_order?: 'asc' | 'desc' + // 错误请求 tab 专属筛选(仅传给错误列表接口;共用同一 filters 对象) + error_phase?: string | null + error_category?: string | null + status_code?: number | null } // ==================== API Functions ==================== diff --git a/frontend/src/assets/icons/payment.svg b/frontend/src/assets/icons/payment.svg new file mode 100644 index 0000000000..c78bea4cf7 --- /dev/null +++ b/frontend/src/assets/icons/payment.svg @@ -0,0 +1,9 @@ + + + + + + + + + diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 67d7b165dc..a4cb70a9c3 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -68,6 +68,15 @@ color="purple" /> + + +
+ +
+
+
+ +

+ {{ t('admin.accounts.headerOverride.hint') }} +

+
+ +
+
+ + +
+
+

+ + {{ t('admin.accounts.headerOverride.info') }} +

+
+ +

+ {{ t('admin.accounts.headerOverride.bulkReplaceHint') }} +

+ +
+
+ + + +
+
+ + + +
+ +
+ +

+ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +

+
+

+ {{ t('admin.accounts.headerOverride.bulkDisableHint') }} +

+
+
+
@@ -1149,6 +1272,16 @@ import { buildModelMappingObject as buildModelMappingPayload, getPresetMappingsByPlatform } from '@/composables/useModelWhitelist' +import { + buildHeaderOverridesObject, + getHeaderOverrideTemplate, + isHeaderOverridePlatform, + validateHeaderOverrideRows, + HEADER_OVERRIDE_ENABLED_CREDENTIAL_KEY, + HEADER_OVERRIDES_CREDENTIAL_KEY, + type HeaderOverrideRow +} from '@/components/account/credentialsBuilder' +import { createStableObjectKeyResolver } from '@/utils/stableObjectKey' import { OPENAI_WS_MODE_CTX_POOL, OPENAI_WS_MODE_OFF, @@ -1217,6 +1350,16 @@ const allOpenAIAPIKey = computed(() => { ) }) +// 是否全部为 anthropic/openai 平台的 apikey 账号(请求头覆写仅在此条件下显示) +const allHeaderOverrideCapable = computed(() => { + return ( + targetSelectedPlatforms.value.length > 0 && + targetSelectedPlatforms.value.every(p => isHeaderOverridePlatform(p)) && + targetSelectedTypes.value.length > 0 && + targetSelectedTypes.value.every(t => t === 'apikey') + ) +}) + // 是否全部为 Anthropic OAuth/SetupToken(RPM 配置仅在此条件下显示) const allAnthropicOAuthOrSetupToken = computed(() => { return ( @@ -1253,6 +1396,7 @@ const enableBaseUrl = ref(false) const enableModelRestriction = ref(false) const enableCustomErrorCodes = ref(false) const enableInterceptWarmup = ref(false) +const enableHeaderOverride = ref(false) const enableProxy = ref(false) const enableConcurrency = ref(false) const enableLoadFactor = ref(false) @@ -1281,6 +1425,39 @@ const modelMappings = ref([]) const selectedErrorCodes = ref([]) const customErrorCodeInput = ref(null) const interceptWarmupRequests = ref(false) +const headerOverrideEnabled = ref(false) +const headerOverrideRows = ref([]) +const getHeaderOverrideRowKey = createStableObjectKeyResolver('bulk-header-override-row') + +const addHeaderOverrideRow = () => { + headerOverrideRows.value.push({ name: '', value: '' }) +} + +const removeHeaderOverrideRow = (index: number) => { + headerOverrideRows.value.splice(index, 1) +} + +// 模板仅在所选账号平台唯一时可用:混合 anthropic+openai 选择无法确定用哪套模板, +// 误填会把另一平台的专有头写进所有所选账号 +const headerOverrideTemplatePlatform = computed(() => { + return targetSelectedPlatforms.value.length === 1 ? targetSelectedPlatforms.value[0] : null +}) + +// 模板按钮:填入所选平台的标准客户端请求头名称(值留空),跳过已存在的同名行 +const fillHeaderOverrideTemplate = () => { + const platform = headerOverrideTemplatePlatform.value + if (!platform) return + const existing = new Set( + headerOverrideRows.value.map((row) => row.name.trim().toLowerCase()).filter(Boolean) + ) + const rows = headerOverrideRows.value.filter((row) => row.name.trim() || row.value.trim()) + for (const row of getHeaderOverrideTemplate(platform)) { + if (!existing.has(row.name)) { + rows.push(row) + } + } + headerOverrideRows.value = rows +} const proxyId = ref(null) const concurrency = ref(1) const loadFactor = ref(null) @@ -1523,6 +1700,15 @@ const buildUpdatePayload = (): Record | null => { credentialsChanged = true } + if (enableHeaderOverride.value) { + // 后端使用 JSONB || merge 语义:关闭时显式写入 false + 空对象以清除旧配置 + credentials[HEADER_OVERRIDE_ENABLED_CREDENTIAL_KEY] = headerOverrideEnabled.value + credentials[HEADER_OVERRIDES_CREDENTIAL_KEY] = headerOverrideEnabled.value + ? buildHeaderOverridesObject(headerOverrideRows.value) + : {} + credentialsChanged = true + } + if (enableOpenAIWSMode.value) { const extra = ensureExtra() extra.openai_oauth_responses_websockets_v2_mode = openaiOAuthResponsesWebSocketV2Mode.value @@ -1651,6 +1837,7 @@ const handleSubmit = async () => { enableModelRestriction.value || enableCustomErrorCodes.value || enableInterceptWarmup.value || + enableHeaderOverride.value || enableProxy.value || enableConcurrency.value || enableLoadFactor.value || @@ -1672,6 +1859,20 @@ const handleSubmit = async () => { return } + if (enableHeaderOverride.value && headerOverrideEnabled.value) { + // 批量保存对 header_overrides 是整键替换:开启但没有任何有效行会把所选账号的 + // 既有覆写配置静默清空,必须显式拦截(清空请走关闭开关的路径,有专门提示) + if (!headerOverrideRows.value.some((row) => row.name.trim())) { + appStore.showError(t('admin.accounts.headerOverride.bulkEmptyRows')) + return + } + const headerError = validateHeaderOverrideRows(headerOverrideRows.value) + if (headerError) { + appStore.showError(t(`admin.accounts.headerOverride.${headerError}`)) + return + } + } + const built = buildUpdatePayload() if (!built) { appStore.showError(t('admin.accounts.bulkEdit.noFieldsSelected')) @@ -1753,6 +1954,7 @@ watch( enableModelRestriction.value = false enableCustomErrorCodes.value = false enableInterceptWarmup.value = false + enableHeaderOverride.value = false enableProxy.value = false enableConcurrency.value = false enableLoadFactor.value = false @@ -1778,6 +1980,8 @@ watch( selectedErrorCodes.value = [] customErrorCodeInput.value = null interceptWarmupRequests.value = false + headerOverrideEnabled.value = false + headerOverrideRows.value = [] proxyId.value = null concurrency.value = 1 loadFactor.value = null diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 3514153c82..67750e4e34 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -1468,6 +1468,110 @@
+ +
+
+
+ +

+ {{ t('admin.accounts.headerOverride.hint') }} +

+
+ +
+ +
+
+

+ + {{ t('admin.accounts.headerOverride.info') }} +

+
+ +
+
+ + + +
+
+ + + +
+ +
+ +

+ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +

+
+
+ @@ -3328,7 +3432,12 @@ import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector. import QuotaLimitCard from '@/components/account/QuotaLimitCard.vue' import { applyAntigravityProjectID, - applyInterceptWarmup + applyHeaderOverride, + applyInterceptWarmup, + getHeaderOverrideTemplate, + isHeaderOverridePlatform, + validateHeaderOverrideRows, + type HeaderOverrideRow } from '@/components/account/credentialsBuilder' import { formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format' import { createStableObjectKeyResolver } from '@/utils/stableObjectKey' @@ -3512,6 +3621,30 @@ function parsePoolModeRetryStatusCodes(input: string): number[] { const customErrorCodesEnabled = ref(false) const selectedErrorCodes = ref([]) const customErrorCodeInput = ref(null) +const headerOverrideEnabled = ref(false) +const headerOverrideRows = ref([]) + +const addHeaderOverrideRow = () => { + headerOverrideRows.value.push({ name: '', value: '' }) +} + +const removeHeaderOverrideRow = (index: number) => { + headerOverrideRows.value.splice(index, 1) +} + +// 模板按钮:填入标准客户端请求头名称(值留空),跳过已存在的同名行 +const fillHeaderOverrideTemplate = () => { + const existing = new Set( + headerOverrideRows.value.map((row) => row.name.trim().toLowerCase()).filter(Boolean) + ) + const rows = headerOverrideRows.value.filter((row) => row.name.trim() || row.value.trim()) + for (const row of getHeaderOverrideTemplate(form.platform)) { + if (!existing.has(row.name)) { + rows.push(row) + } + } + headerOverrideRows.value = rows +} const interceptWarmupRequests = ref(false) const autoPauseOnExpired = ref(true) const openaiPassthroughEnabled = ref(false) @@ -3569,6 +3702,7 @@ const vertexServiceAccountDragActive = ref(false) const tempUnschedEnabled = ref(false) const tempUnschedRules = ref([]) const getModelMappingKey = createStableObjectKeyResolver('create-model-mapping') +const getHeaderOverrideRowKey = createStableObjectKeyResolver('create-header-override-row') const getOpenAICompactModelMappingKey = createStableObjectKeyResolver('create-openai-compact-model-mapping') const getAntigravityModelMappingKey = createStableObjectKeyResolver('create-antigravity-model-mapping') const getTempUnschedRuleKey = createStableObjectKeyResolver('create-temp-unsched-rule') @@ -3970,6 +4104,10 @@ watch( anthropicAPIKeyAuthScheme.value = 'x_api_key' webSearchEmulationMode.value = 'default' } + // 请求头覆写为平台相关配置(模板/常用头集合不同),切换平台时清空, + // 避免上一平台的模板行被提交到新平台账号 + headerOverrideEnabled.value = false + headerOverrideRows.value = [] // Reset OAuth states oauth.resetState() openaiOAuth.resetState() @@ -4359,6 +4497,8 @@ const resetForm = () => { customErrorCodesEnabled.value = false selectedErrorCodes.value = [] customErrorCodeInput.value = null + headerOverrideEnabled.value = false + headerOverrideRows.value = [] interceptWarmupRequests.value = false autoPauseOnExpired.value = true openaiPassthroughEnabled.value = false @@ -4789,6 +4929,18 @@ const handleSubmit = async () => { credentials.custom_error_codes = [...selectedErrorCodes.value] } + // Add header override if enabled (anthropic/openai apikey only) + if (isHeaderOverridePlatform(form.platform)) { + if (headerOverrideEnabled.value) { + const headerError = validateHeaderOverrideRows(headerOverrideRows.value) + if (headerError) { + appStore.showError(t(`admin.accounts.headerOverride.${headerError}`)) + return + } + } + applyHeaderOverride(credentials, headerOverrideEnabled.value, headerOverrideRows.value, 'create') + } + applyInterceptWarmup(credentials, interceptWarmupRequests.value, 'create') if (!applyTempUnschedConfig(credentials)) { return diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index b5f97753d8..9b5ab82fc5 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -417,11 +417,115 @@ + +
+
+
+ +

+ {{ t('admin.accounts.headerOverride.hint') }} +

+
+ +
+ +
+
+

+ + {{ t('admin.accounts.headerOverride.info') }} +

+
+ +
+
+ + + +
+
+ + + +
+ +
+ +

+ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +

+
+
+ - +
@@ -1370,7 +1474,7 @@
- +
- + - {{ codexImageGenerationBridgeBadgeLabel }} + {{ codexImageToolBadgeLabel }}

- {{ t('admin.accounts.openai.codexImageGenerationBridgeDesc') }} + {{ t('admin.accounts.openai.codexImageToolDesc') }}

-
+
+
+ +
-
+
+ + +

+ {{ t('payment.admin.subscriptionCnyPayPreview', { amount: subscriptionCnyPreview.amount }) }} + + {{ t('payment.admin.subscriptionCnyPayPreviewWithFee', { feeRate: subscriptionCnyPreview.feeRate, total: subscriptionCnyPreview.total }) }} + +

+
@@ -81,7 +90,9 @@ import { ref, reactive, computed, watch } from 'vue' import { useI18n } from 'vue-i18n' import { useAppStore } from '@/stores/app' import { adminPaymentAPI } from '@/api/admin/payment' +import type { AdminPaymentConfig } from '@/api/admin/payment' import { extractApiErrorMessage } from '@/utils/apiError' +import { formatPaymentAmount } from '@/components/payment/currency' import type { SubscriptionPlan } from '@/types/payment' import type { AdminGroup } from '@/types' import BaseDialog from '@/components/common/BaseDialog.vue' @@ -94,6 +105,7 @@ const props = defineProps<{ show: boolean plan: SubscriptionPlan | null groups: AdminGroup[] + paymentConfig?: AdminPaymentConfig | null }>() const emit = defineEmits<{ @@ -129,6 +141,31 @@ const selectedGroupInfo = computed(() => { return props.groups.find(g => g.id === planForm.group_id) || null }) +function roundCnyAmount(value: number): number { + return Math.round(value * 100) / 100 +} + +function ceilCnyAmount(value: number): number { + return Math.ceil(value * 100) / 100 +} + +const subscriptionCnyPreview = computed(() => { + const price = Number(planForm.price) || 0 + const rate = Number(props.paymentConfig?.subscription_usd_to_cny_rate) || 0 + if (price <= 0 || rate <= 0) return null + + const amount = roundCnyAmount(price * rate) + const feeRate = Number(props.paymentConfig?.recharge_fee_rate) || 0 + const fee = feeRate > 0 ? ceilCnyAmount((amount * feeRate) / 100) : 0 + const total = feeRate > 0 ? roundCnyAmount(amount + fee) : amount + + return { + amount: formatPaymentAmount(amount, 'CNY'), + feeRate, + total: formatPaymentAmount(total, 'CNY'), + } +}) + // Reset form when dialog opens watch(() => props.show, (visible) => { if (!visible) return diff --git a/frontend/src/views/admin/orders/__tests__/PlanEditDialog.spec.ts b/frontend/src/views/admin/orders/__tests__/PlanEditDialog.spec.ts new file mode 100644 index 0000000000..9c31e7c176 --- /dev/null +++ b/frontend/src/views/admin/orders/__tests__/PlanEditDialog.spec.ts @@ -0,0 +1,77 @@ +import { describe, expect, it, vi } from 'vitest' +import { mount } from '@vue/test-utils' +import PlanEditDialog from '../PlanEditDialog.vue' + +vi.mock('vue-i18n', () => ({ + useI18n: () => ({ + t: (key: string, params?: Record) => { + if (key === 'payment.admin.subscriptionCnyPayPreview') return `preview ${params?.amount}` + if (key === 'payment.admin.subscriptionCnyPayPreviewWithFee') return `fee ${params?.feeRate} ${params?.total}` + return key + }, + }), +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => ({ + showError: vi.fn(), + showSuccess: vi.fn(), + }), +})) + +vi.mock('@/api/admin/payment', () => ({ + adminPaymentAPI: { + createPlan: vi.fn(), + updatePlan: vi.fn(), + }, +})) + +function mountDialog(paymentConfig: Record | null) { + return mount(PlanEditDialog, { + props: { + show: true, + plan: null, + groups: [], + paymentConfig, + }, + global: { + stubs: { + BaseDialog: { + props: ['show'], + template: '
', + }, + Select: true, + Icon: true, + GroupBadge: true, + }, + }, + }) +} + +describe('PlanEditDialog subscription CNY payment preview', () => { + it('shows CNY channel charge using the configured subscription rate and fee', async () => { + const wrapper = mountDialog({ + subscription_usd_to_cny_rate: 7.15, + recharge_fee_rate: 2.5, + }) + + await wrapper.find('input[type="number"]').setValue('9.99') + + expect(wrapper.text()).toContain('preview') + expect(wrapper.text()).toContain('¥71.43') + expect(wrapper.text()).toContain('fee 2.5') + expect(wrapper.text()).toContain('¥73.22') + }) + + it('hides the preview when the subscription rate is not configured', async () => { + const wrapper = mountDialog({ + subscription_usd_to_cny_rate: 0, + recharge_fee_rate: 2.5, + }) + + await wrapper.find('input[type="number"]').setValue('9.99') + + expect(wrapper.text()).not.toContain('preview') + expect(wrapper.text()).not.toContain('¥71.43') + }) +}) diff --git a/frontend/src/views/user/KeysView.vue b/frontend/src/views/user/KeysView.vue index 1aabf6b834..087e0e4175 100644 --- a/frontend/src/views/user/KeysView.vue +++ b/frontend/src/views/user/KeysView.vue @@ -170,6 +170,19 @@
+ +
proxy4freeProxy4Free は開発者と AI アプリケーション向けのデータプロキシサービスプロバイダーで、住宅プロキシ、静的住宅プロキシ、ISP プロキシ、データセンタープロキシなど多様なプロキシソリューションを提供しており、Web Scraping、Browser Automation、AI Agent などのシナリオに適しています。グローバル IP リソース、安定した接続、柔軟な切り替えをサポートし、開発者のデータ収集成功率の向上と IP ブロックリスクの低減を支援します。こちらのリンクから登録して、より安定した効率的な自動化ワークフローを簡単に構築しましょう。 +Proxy4Free のご支援に感謝します!Proxy4Free は開発者と AI アプリケーション向けのデータプロキシサービスプロバイダーで、住宅プロキシ、静的住宅プロキシ、ISP プロキシ、データセンタープロキシなど多様なプロキシソリューションを提供しており、Web Scraping、Browser Automation、AI Agent などのシナリオに適しています。グローバル IP リソース、安定した接続、柔軟な切り替えをサポートし、開発者のデータ収集成功率の向上と IP ブロックリスクの低減を支援します。こちらのリンクから登録して、より安定した効率的な自動化ワークフローを簡単に構築しましょう。