diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml index bb5cf692bc..8599465254 100644 --- a/.github/workflows/backend-ci.yml +++ b/.github/workflows/backend-ci.yml @@ -8,6 +8,15 @@ permissions: contents: read jobs: + shell: + runs-on: macos-15 + steps: + - uses: actions/checkout@v6 + - name: Check deployment scripts + run: | + /bin/bash -n deploy/apple-container.sh + /bin/bash deploy/tests/apple-container-test.sh + test: runs-on: ubuntu-latest steps: diff --git a/.gitignore b/.gitignore index bd2e3e6ddf..d45d27f79e 100644 --- a/.gitignore +++ b/.gitignore @@ -116,6 +116,8 @@ backend/.installed # 其他 # =================== tests +!deploy/tests/ +!deploy/tests/** CLAUDE.md .claude scripts diff --git a/README.md b/README.md index 661cd8e7eb..00f021ea2c 100644 --- a/README.md +++ b/README.md @@ -163,6 +163,12 @@ Model authenticity: no content intervention or secondary filtering — experienc + +aimzoon +Thanks to Aimzoon for sponsoring this project! Aimzoon provides stable, cost-effective AI API access services, enabling developers to quickly connect popular AI services to coding tools such as Codex, Claude Code, and Gemini CLI. No complex configuration — faster onboarding, more stable calls, and lower costs. Ongoing promotions including discounted Codex rates and special pricing, with free trial credits upon registration, bringing AI coding into your daily workflow. Click here to register and try it out! + + + ## Overview @@ -329,6 +335,7 @@ cd sub2api/deploy # 2. Copy environment configuration cp .env.example .env +chmod 600 .env # 3. Edit configuration (generate secure passwords) nano .env @@ -448,7 +455,23 @@ rm -rf data/ postgres_data/ redis_data/ --- -### Method 3: Build from Source +### Method 3: Apple container (macOS) + +Apple-silicon Macs running macOS 26 can run the full Sub2API, PostgreSQL, and Redis stack with Apple `container` 1.1.0 or newer: + +```bash +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api/deploy +./apple-container.sh init +./apple-container.sh up +./apple-container.sh status +``` + +This is an operator-managed local workflow; Docker Compose remains the recommended production path. See [deploy/APPLE_CONTAINER.md](deploy/APPLE_CONTAINER.md) for lifecycle commands, persistence, upgrades, and runtime limitations. + +--- + +### Method 4: Build from Source Build and run from source code for development or customization. @@ -579,6 +602,27 @@ If you disable URL validation or response header filtering, harden your network - Enforce TLS-only outbound traffic - Strip sensitive upstream response headers at the proxy +#### OpenAI Responses WebSocket ingress limits + +`gateway.openai_ws` bounds the lifetime and aggregate count of client-facing +Responses WebSocket sessions. These safeguards apply independently from +per-turn user and account concurrency slots, which are released between turns. + +```yaml +gateway: + openai_ws: + # Close a client socket idle between completed turns; 0 disables this safeguard. + ingress_inter_turn_idle_timeout_seconds: 300 + # Distributed API-key limit for live client ingress sessions; 0 disables it. + max_ingress_connections_per_api_key: 64 +``` + +The connection cap is coordinated through Redis using a 60-second lease that +is refreshed every 20 seconds. A process that cannot confirm a lease for a +full lease lifetime closes its local WebSocket rather than continuing outside +the global cap. Use `http_bridge` for client-WebSocket/upstream-HTTP operation +when rolling out or mitigating upstream WebSocket issues. + #### ⚠️ Important: Creating the Admin Account The initial admin account is **only created via the setup wizard** (served at `http://:8080` on first run). The `default.admin_email` / `default.admin_password` fields in `config.yaml` are **not used** to create it — they exist in the template for historical reasons. @@ -637,20 +681,20 @@ Simple Mode is designed for individual developers or internal teams who want qui --- -## Grok / xAI OAuth Support +## Grok / xAI Support -Sub2API supports Grok subscription accounts through xAI OAuth and forwards OpenAI-compatible Responses traffic to xAI. +Sub2API supports both Grok subscription accounts through xAI OAuth and standard xAI API-key accounts. Both account types forward OpenAI-compatible Responses traffic to xAI. ### Supported Scope - Platform name: `grok` -- Account type: OAuth subscription accounts -- Public Responses targets: `/v1/responses`, `/responses`, and `/backend-api/codex/responses`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/responses` +- Account types: OAuth subscription accounts and xAI API-key accounts +- Public Responses targets: `/v1/responses`, `/responses`, and `/backend-api/codex/responses`, forwarded to the Grok subscription proxy for OAuth accounts or `https://api.x.ai/v1/responses` for API-key accounts - 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` +- Public Chat Completions targets: `/v1/chat/completions` and `/chat/completions`, forwarded to the account-type-specific xAI upstream - Codex CLI style Responses WebSocket ingress is accepted on the Responses targets and bridged to xAI HTTP/SSE Responses upstream -- 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. +- Text models: `grok-4.5`, `grok-4.3`, `grok-build-0.1`, `grok-composer-2.5-fast`, `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/edits`, `/videos/edits`, `/v1/videos/extensions`, `/videos/extensions`, `/v1/videos/{request_id}`, and `/videos/{request_id}`. Generation, editing, and extension 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 @@ -665,9 +709,10 @@ The Grok OAuth flow uses PKCE and does not require committing private secrets. T | `XAI_OAUTH_REDIRECT_URI` | `http://127.0.0.1:56121/callback` | | `XAI_OAUTH_AUTHORIZE_URL` | `https://auth.x.ai/oauth2/authorize` | | `XAI_OAUTH_TOKEN_URL` | `https://auth.x.ai/oauth2/token` | -| `XAI_BASE_URL` | `https://api.x.ai/v1` | +| `XAI_BASE_URL` | `https://api.x.ai/v1`; runtime-diagnostics override (account `base_url` controls request forwarding) | +| `XAI_GROK_CLI_VERSION` | `0.2.93`; optional override for the client identity sent to `cli-chat-proxy.grok.com` | -Administrators can create or reauthorize Grok accounts from the dashboard, or use the admin API: +Administrators can create Grok OAuth or API-key accounts from the dashboard. OAuth authorization and reauthorization are also available through the admin API: | Endpoint | Purpose | |----------|---------| @@ -676,13 +721,47 @@ Administrators can create or reauthorize Grok accounts from the dashboard, or us | `POST /api/v1/admin/grok/oauth/refresh-token` | Validate or refresh a Grok refresh token | | `POST /api/v1/admin/grok/accounts/:id/refresh` | Refresh an existing Grok account | -Credential storage reuses the existing account JSON fields: `access_token`, `refresh_token`, `token_type`, `expires_at`, optional `email`, optional `subscription_tier`, and `entitlement_status`. +OAuth credential storage reuses the existing account JSON fields: `access_token`, `refresh_token`, `token_type`, `expires_at`, `base_url`, optional `email`, optional `subscription_tier`, and `entitlement_status`. OAuth inference defaults to `https://cli-chat-proxy.grok.com/v1`; existing OAuth accounts that stored the old `https://api.x.ai/v1` default are redirected to the subscription proxy at runtime. Explicit custom upstreams remain unchanged. + +For API-key accounts, select **Grok → API Key** in the create-account dialog. The official base URL defaults to `https://api.x.ai/v1`; credentials use the existing `base_url` and `api_key` account fields. OAuth accounts continue to use the subscription flow above. + +### Grok Build CLI Configuration + +1. In the Sub2API admin dashboard, add either a `grok` OAuth account and complete xAI authorization, or add a Grok API-key account. +2. Create a Grok group, attach the account to it, then create a Sub2API API key assigned to that group. +3. In the user API-key page, click **Use Key** and select **Grok CLI**. The modal generates the correct file and base URL for macOS/Linux or Windows. It also provides an OpenCode configuration on the **OpenCode** tab. +4. If configuring manually, save the following as `~/.grok/config.toml` (Windows: `%USERPROFILE%\.grok\config.toml`): + +```toml +[models] +default = "sub2api-grok" +web_search = "sub2api-grok" + +[model."sub2api-grok"] +model = "grok-4.5" +base_url = "https://your-sub2api.example.com/v1" +name = "Grok 4.5 via Sub2API" +description = "Grok 4.5 through a Sub2API Grok group" +api_key = "sk-your-sub2api-key" +api_backend = "responses" +context_window = 1000000 +supports_backend_search = true +``` + +Back up an existing `config.toml` before merging the entry. The file contains a Sub2API API key, so keep it private and restrict its permissions where supported. Verify the effective configuration and make a smoke request: + +```bash +grok inspect +grok -p "Reply with sub2api-ok" -m sub2api-grok +``` + +The `base_url` above is the public Sub2API URL ending in `/v1`, not `api.x.ai` or the internal xAI OAuth proxy URL. ### Usage And Quota Display xAI quota is passive. Sub2API does not invent subscription quota values; it records whitelisted xAI rate-limit headers from successful or rate-limited upstream responses when xAI sends them. Before the first usable upstream response, the dashboard shows quota as unknown and still displays local Sub2API usage stats. -`401` responses mark the account as needing reauthorization. `403` responses are treated as entitlement or subscription-tier failures instead of token-refresh loops. `429` responses use `Retry-After` or a short cooldown to temporarily remove the account from scheduling. +`401` responses temporarily remove accounts with invalid credentials from scheduling. `403` responses are treated as access or entitlement failures instead of token-refresh loops. `429` responses use `Retry-After` or a short cooldown to temporarily remove the account from scheduling. --- diff --git a/README_CN.md b/README_CN.md index 88c8ce11b1..b1139dd122 100644 --- a/README_CN.md +++ b/README_CN.md @@ -166,6 +166,11 @@ + +aimzoon +感谢 Aimzoon 对本项目的赞助! Aimzoon 提供稳定、高性价比的 AI API 接入服务,支持开发者将常用 AI 服务快速接入 Codex、Claude Code、Gemini CLI 等编程工具。无需复杂配置,更快接入,更稳调用,更省成本。codex倍率优惠,特价倍率等促销不断,注册即送免费体验额度,让 AI 编程真正进入日常工作流。点击这里注册体验! + + @@ -333,6 +338,7 @@ cd sub2api/deploy # 2. 复制环境配置文件 cp .env.example .env +chmod 600 .env # 3. 编辑配置(生成安全密码) nano .env @@ -464,7 +470,23 @@ rm -rf data/ postgres_data/ redis_data/ --- -### 方式三:源码编译 +### 方式三:Apple container(macOS) + +Apple 芯片 Mac 在 macOS 26 上可使用 Apple `container` 1.1.0 或更高版本运行完整的 Sub2API、PostgreSQL 和 Redis: + +```bash +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api/deploy +./apple-container.sh init +./apple-container.sh up +./apple-container.sh status +``` + +该方式面向本地开发和人工运维,不提供持续重启监管;生产部署仍推荐 Docker Compose。生命周期命令、持久化、升级和运行时限制见 [deploy/APPLE_CONTAINER.md](deploy/APPLE_CONTAINER.md)。 + +--- + +### 方式四:源码编译 从源码编译安装,适合开发或定制需求。 diff --git a/README_JA.md b/README_JA.md index 21d070e397..ba556139ee 100644 --- a/README_JA.md +++ b/README_JA.md @@ -161,6 +161,12 @@ + +aimzoon +Aimzoon のご支援に感謝します!Aimzoon は安定してコストパフォーマンスに優れた AI API 接続サービスを提供し、開発者が主要な AI サービスを Codex、Claude Code、Gemini CLI などのコーディングツールへ素早く接続できるようにします。複雑な設定は不要で、より速い接続、より安定した呼び出し、より低いコストを実現。Codex レート割引や特価レートなどのキャンペーンも随時開催中、登録するだけで無料お試しクレジットをプレゼント。AI コーディングを日常のワークフローへ。こちらから登録してお試しください! + + + ## 概要 @@ -327,6 +333,7 @@ cd sub2api/deploy # 2. 環境設定ファイルをコピー cp .env.example .env +chmod 600 .env # 3. 設定を編集(セキュアなパスワードを生成) nano .env @@ -446,7 +453,23 @@ rm -rf data/ postgres_data/ redis_data/ --- -### 方法3: ソースからビルド +### 方法3: Apple container(macOS) + +Apple シリコン搭載 Mac と macOS 26 では、Apple `container` 1.1.0 以降を使用して Sub2API、PostgreSQL、Redis の完全なスタックを実行できます: + +```bash +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api/deploy +./apple-container.sh init +./apple-container.sh up +./apple-container.sh status +``` + +これはローカル開発および手動運用向けです。本番環境では引き続き Docker Compose を推奨します。ライフサイクル、永続化、アップグレード、制限については [deploy/APPLE_CONTAINER.md](deploy/APPLE_CONTAINER.md) を参照してください。 + +--- + +### 方法4: ソースからビルド 開発やカスタマイズのためにソースコードからビルドして実行します。 diff --git a/assets/partners/logos/aimzoon.jpg b/assets/partners/logos/aimzoon.jpg new file mode 100644 index 0000000000..3fa8c653c8 Binary files /dev/null and b/assets/partners/logos/aimzoon.jpg differ diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 23c64fda6f..7a29ae6cd9 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.150 +0.1.155 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 088046de96..2b339f3d96 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -182,12 +182,13 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream) antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository) grokQuotaFetcher := service.NewGrokQuotaFetcher() + grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, usageLogRepository) openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory, openAIGatewayService) usageCache := service.NewUsageCache() - accountUsageService := service.ProvideAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService, openAIGatewayService) + accountUsageService := service.ProvideAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService, openAIGatewayService) accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService) crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig) - accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator) + accountHandler := admin.ProvideAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, grokQuotaService) adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService) dataManagementService := service.NewDataManagementService() dataManagementHandler := admin.NewDataManagementHandler(dataManagementService) @@ -199,7 +200,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIOAuthHandler := admin.NewOpenAIOAuthHandler(openAIOAuthService, adminService, openAIQuotaService) geminiOAuthHandler := admin.NewGeminiOAuthHandler(geminiOAuthService) antigravityOAuthHandler := admin.NewAntigravityOAuthHandler(antigravityOAuthService) - grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream) grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService) proxyHandler := admin.NewProxyHandler(adminService) adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService) @@ -256,7 +256,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIGatewayHandler := handler.NewOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig) handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService) totpHandler := handler.NewTotpHandler(totpService) - handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService, channelService) + handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService) paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry) availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService) batchImageRepository := repository.NewBatchImageRepository(db) diff --git a/backend/ent/channelmonitor/channelmonitor.go b/backend/ent/channelmonitor/channelmonitor.go index afdc6957d6..711e6217e2 100644 --- a/backend/ent/channelmonitor/channelmonitor.go +++ b/backend/ent/channelmonitor/channelmonitor.go @@ -167,6 +167,7 @@ const ( ProviderOpenai Provider = "openai" ProviderAnthropic Provider = "anthropic" ProviderGemini Provider = "gemini" + ProviderGrok Provider = "grok" ) func (pr Provider) String() string { @@ -176,7 +177,7 @@ func (pr Provider) String() string { // ProviderValidator is a validator for the "provider" field enum values. It is called by the builders before save. func ProviderValidator(pr Provider) error { switch pr { - case ProviderOpenai, ProviderAnthropic, ProviderGemini: + case ProviderOpenai, ProviderAnthropic, ProviderGemini, ProviderGrok: return nil default: return fmt.Errorf("channelmonitor: invalid enum value for provider field: %q", pr) diff --git a/backend/ent/channelmonitorrequesttemplate/channelmonitorrequesttemplate.go b/backend/ent/channelmonitorrequesttemplate/channelmonitorrequesttemplate.go index db04aee106..5989d0e743 100644 --- a/backend/ent/channelmonitorrequesttemplate/channelmonitorrequesttemplate.go +++ b/backend/ent/channelmonitorrequesttemplate/channelmonitorrequesttemplate.go @@ -103,6 +103,7 @@ const ( ProviderOpenai Provider = "openai" ProviderAnthropic Provider = "anthropic" ProviderGemini Provider = "gemini" + ProviderGrok Provider = "grok" ) func (pr Provider) String() string { @@ -112,7 +113,7 @@ func (pr Provider) String() string { // ProviderValidator is a validator for the "provider" field enum values. It is called by the builders before save. func ProviderValidator(pr Provider) error { switch pr { - case ProviderOpenai, ProviderAnthropic, ProviderGemini: + case ProviderOpenai, ProviderAnthropic, ProviderGemini, ProviderGrok: return nil default: return fmt.Errorf("channelmonitorrequesttemplate: invalid enum value for provider field: %q", pr) diff --git a/backend/ent/group.go b/backend/ent/group.go index 5bec594977..088da069d6 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -83,6 +83,8 @@ type Group struct { VideoPrice720p *float64 `json:"video_price_720p,omitempty"` // VideoPrice1080p holds the value of the "video_price_1080p" field. VideoPrice1080p *float64 `json:"video_price_1080p,omitempty"` + // Codex alpha/search 网页搜索单次价格(USD/次);nil 表示使用默认价 0.01(官方 $10/1000 次) + WebSearchPricePerCall *float64 `json:"web_search_price_per_call,omitempty"` // 是否仅允许 Claude Code 客户端 ClaudeCodeOnly bool `json:"claude_code_only,omitempty"` // 非 Claude Code 请求降级使用的分组 ID @@ -223,7 +225,7 @@ func (*Group) scanValues(columns []string) ([]any, error) { values[i] = new([]byte) case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: values[i] = new(sql.NullBool) - case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p: + case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall: values[i] = new(sql.NullFloat64) case group.FieldID, group.FieldDefaultValidityDays, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, group.FieldSortOrder, group.FieldRpmLimit: values[i] = new(sql.NullInt64) @@ -455,6 +457,13 @@ func (_m *Group) assignValues(columns []string, values []any) error { _m.VideoPrice1080p = new(float64) *_m.VideoPrice1080p = value.Float64 } + case group.FieldWebSearchPricePerCall: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field web_search_price_per_call", values[i]) + } else if value.Valid { + _m.WebSearchPricePerCall = new(float64) + *_m.WebSearchPricePerCall = value.Float64 + } case group.FieldClaudeCodeOnly: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field claude_code_only", values[i]) @@ -749,6 +758,11 @@ func (_m *Group) String() string { builder.WriteString(fmt.Sprintf("%v", *v)) } builder.WriteString(", ") + if v := _m.WebSearchPricePerCall; v != nil { + builder.WriteString("web_search_price_per_call=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") builder.WriteString("claude_code_only=") builder.WriteString(fmt.Sprintf("%v", _m.ClaudeCodeOnly)) builder.WriteString(", ") diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 769c63e6b1..61d7a21d67 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -80,6 +80,8 @@ const ( FieldVideoPrice720p = "video_price_720p" // FieldVideoPrice1080p holds the string denoting the video_price_1080p field in the database. FieldVideoPrice1080p = "video_price_1080p" + // FieldWebSearchPricePerCall holds the string denoting the web_search_price_per_call field in the database. + FieldWebSearchPricePerCall = "web_search_price_per_call" // FieldClaudeCodeOnly holds the string denoting the claude_code_only field in the database. FieldClaudeCodeOnly = "claude_code_only" // FieldFallbackGroupID holds the string denoting the fallback_group_id field in the database. @@ -217,6 +219,7 @@ var Columns = []string{ FieldVideoPrice480p, FieldVideoPrice720p, FieldVideoPrice1080p, + FieldWebSearchPricePerCall, FieldClaudeCodeOnly, FieldFallbackGroupID, FieldFallbackGroupIDOnInvalidRequest, @@ -511,6 +514,11 @@ func ByVideoPrice1080p(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldVideoPrice1080p, opts...).ToFunc() } +// ByWebSearchPricePerCall orders the results by the web_search_price_per_call field. +func ByWebSearchPricePerCall(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldWebSearchPricePerCall, opts...).ToFunc() +} + // ByClaudeCodeOnly orders the results by the claude_code_only field. func ByClaudeCodeOnly(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldClaudeCodeOnly, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 5a9d92d0f4..b9a52a2eb6 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -215,6 +215,11 @@ func VideoPrice1080p(v float64) predicate.Group { return predicate.Group(sql.FieldEQ(FieldVideoPrice1080p, v)) } +// WebSearchPricePerCall applies equality check predicate on the "web_search_price_per_call" field. It's identical to WebSearchPricePerCallEQ. +func WebSearchPricePerCall(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldWebSearchPricePerCall, v)) +} + // ClaudeCodeOnly applies equality check predicate on the "claude_code_only" field. It's identical to ClaudeCodeOnlyEQ. func ClaudeCodeOnly(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) @@ -1655,6 +1660,56 @@ func VideoPrice1080pNotNil() predicate.Group { return predicate.Group(sql.FieldNotNull(FieldVideoPrice1080p)) } +// WebSearchPricePerCallEQ applies the EQ predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallNEQ applies the NEQ predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallIn applies the In predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldWebSearchPricePerCall, vs...)) +} + +// WebSearchPricePerCallNotIn applies the NotIn predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldWebSearchPricePerCall, vs...)) +} + +// WebSearchPricePerCallGT applies the GT predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallGTE applies the GTE predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallLT applies the LT predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallLTE applies the LTE predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldWebSearchPricePerCall, v)) +} + +// WebSearchPricePerCallIsNil applies the IsNil predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldWebSearchPricePerCall)) +} + +// WebSearchPricePerCallNotNil applies the NotNil predicate on the "web_search_price_per_call" field. +func WebSearchPricePerCallNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldWebSearchPricePerCall)) +} + // ClaudeCodeOnlyEQ applies the EQ predicate on the "claude_code_only" field. func ClaudeCodeOnlyEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index 2a6c18e67d..53fd733c1f 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -469,6 +469,20 @@ func (_c *GroupCreate) SetNillableVideoPrice1080p(v *float64) *GroupCreate { return _c } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (_c *GroupCreate) SetWebSearchPricePerCall(v float64) *GroupCreate { + _c.mutation.SetWebSearchPricePerCall(v) + return _c +} + +// SetNillableWebSearchPricePerCall sets the "web_search_price_per_call" field if the given value is not nil. +func (_c *GroupCreate) SetNillableWebSearchPricePerCall(v *float64) *GroupCreate { + if v != nil { + _c.SetWebSearchPricePerCall(*v) + } + return _c +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_c *GroupCreate) SetClaudeCodeOnly(v bool) *GroupCreate { _c.mutation.SetClaudeCodeOnly(v) @@ -1218,6 +1232,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldVideoPrice1080p, field.TypeFloat64, value) _node.VideoPrice1080p = &value } + if value, ok := _c.mutation.WebSearchPricePerCall(); ok { + _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) + _node.WebSearchPricePerCall = &value + } if value, ok := _c.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) _node.ClaudeCodeOnly = value @@ -1968,6 +1986,30 @@ func (u *GroupUpsert) ClearVideoPrice1080p() *GroupUpsert { return u } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (u *GroupUpsert) SetWebSearchPricePerCall(v float64) *GroupUpsert { + u.Set(group.FieldWebSearchPricePerCall, v) + return u +} + +// UpdateWebSearchPricePerCall sets the "web_search_price_per_call" field to the value that was provided on create. +func (u *GroupUpsert) UpdateWebSearchPricePerCall() *GroupUpsert { + u.SetExcluded(group.FieldWebSearchPricePerCall) + return u +} + +// AddWebSearchPricePerCall adds v to the "web_search_price_per_call" field. +func (u *GroupUpsert) AddWebSearchPricePerCall(v float64) *GroupUpsert { + u.Add(group.FieldWebSearchPricePerCall, v) + return u +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (u *GroupUpsert) ClearWebSearchPricePerCall() *GroupUpsert { + u.SetNull(group.FieldWebSearchPricePerCall) + return u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsert) SetClaudeCodeOnly(v bool) *GroupUpsert { u.Set(group.FieldClaudeCodeOnly, v) @@ -2858,6 +2900,34 @@ func (u *GroupUpsertOne) ClearVideoPrice1080p() *GroupUpsertOne { }) } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (u *GroupUpsertOne) SetWebSearchPricePerCall(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetWebSearchPricePerCall(v) + }) +} + +// AddWebSearchPricePerCall adds v to the "web_search_price_per_call" field. +func (u *GroupUpsertOne) AddWebSearchPricePerCall(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddWebSearchPricePerCall(v) + }) +} + +// UpdateWebSearchPricePerCall sets the "web_search_price_per_call" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateWebSearchPricePerCall() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateWebSearchPricePerCall() + }) +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (u *GroupUpsertOne) ClearWebSearchPricePerCall() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearWebSearchPricePerCall() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertOne) SetClaudeCodeOnly(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -3951,6 +4021,34 @@ func (u *GroupUpsertBulk) ClearVideoPrice1080p() *GroupUpsertBulk { }) } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (u *GroupUpsertBulk) SetWebSearchPricePerCall(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetWebSearchPricePerCall(v) + }) +} + +// AddWebSearchPricePerCall adds v to the "web_search_price_per_call" field. +func (u *GroupUpsertBulk) AddWebSearchPricePerCall(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddWebSearchPricePerCall(v) + }) +} + +// UpdateWebSearchPricePerCall sets the "web_search_price_per_call" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateWebSearchPricePerCall() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateWebSearchPricePerCall() + }) +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (u *GroupUpsertBulk) ClearWebSearchPricePerCall() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearWebSearchPricePerCall() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertBulk) SetClaudeCodeOnly(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index 3bb18d3e1a..1e767a139b 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -640,6 +640,33 @@ func (_u *GroupUpdate) ClearVideoPrice1080p() *GroupUpdate { return _u } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (_u *GroupUpdate) SetWebSearchPricePerCall(v float64) *GroupUpdate { + _u.mutation.ResetWebSearchPricePerCall() + _u.mutation.SetWebSearchPricePerCall(v) + return _u +} + +// SetNillableWebSearchPricePerCall sets the "web_search_price_per_call" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableWebSearchPricePerCall(v *float64) *GroupUpdate { + if v != nil { + _u.SetWebSearchPricePerCall(*v) + } + return _u +} + +// AddWebSearchPricePerCall adds value to the "web_search_price_per_call" field. +func (_u *GroupUpdate) AddWebSearchPricePerCall(v float64) *GroupUpdate { + _u.mutation.AddWebSearchPricePerCall(v) + return _u +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (_u *GroupUpdate) ClearWebSearchPricePerCall() *GroupUpdate { + _u.mutation.ClearWebSearchPricePerCall() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdate) SetClaudeCodeOnly(v bool) *GroupUpdate { _u.mutation.SetClaudeCodeOnly(v) @@ -1375,6 +1402,15 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.VideoPrice1080pCleared() { _spec.ClearField(group.FieldVideoPrice1080p, field.TypeFloat64) } + if value, ok := _u.mutation.WebSearchPricePerCall(); ok { + _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedWebSearchPricePerCall(); ok { + _spec.AddField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) + } + if _u.mutation.WebSearchPricePerCallCleared() { + _spec.ClearField(group.FieldWebSearchPricePerCall, field.TypeFloat64) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } @@ -2364,6 +2400,33 @@ func (_u *GroupUpdateOne) ClearVideoPrice1080p() *GroupUpdateOne { return _u } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (_u *GroupUpdateOne) SetWebSearchPricePerCall(v float64) *GroupUpdateOne { + _u.mutation.ResetWebSearchPricePerCall() + _u.mutation.SetWebSearchPricePerCall(v) + return _u +} + +// SetNillableWebSearchPricePerCall sets the "web_search_price_per_call" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableWebSearchPricePerCall(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetWebSearchPricePerCall(*v) + } + return _u +} + +// AddWebSearchPricePerCall adds value to the "web_search_price_per_call" field. +func (_u *GroupUpdateOne) AddWebSearchPricePerCall(v float64) *GroupUpdateOne { + _u.mutation.AddWebSearchPricePerCall(v) + return _u +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (_u *GroupUpdateOne) ClearWebSearchPricePerCall() *GroupUpdateOne { + _u.mutation.ClearWebSearchPricePerCall() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdateOne) SetClaudeCodeOnly(v bool) *GroupUpdateOne { _u.mutation.SetClaudeCodeOnly(v) @@ -3129,6 +3192,15 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if _u.mutation.VideoPrice1080pCleared() { _spec.ClearField(group.FieldVideoPrice1080p, field.TypeFloat64) } + if value, ok := _u.mutation.WebSearchPricePerCall(); ok { + _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedWebSearchPricePerCall(); ok { + _spec.AddField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) + } + if _u.mutation.WebSearchPricePerCallCleared() { + _spec.ClearField(group.FieldWebSearchPricePerCall, field.TypeFloat64) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index d3e8bc5448..52229f9151 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -623,7 +623,7 @@ var ( {Name: "created_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "updated_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "name", Type: field.TypeString, Size: 100}, - {Name: "provider", Type: field.TypeEnum, Enums: []string{"openai", "anthropic", "gemini"}}, + {Name: "provider", Type: field.TypeEnum, Enums: []string{"openai", "anthropic", "gemini", "grok"}}, {Name: "api_mode", Type: field.TypeString, Size: 32, Default: "chat_completions"}, {Name: "endpoint", Type: field.TypeString, Size: 500}, {Name: "api_key_encrypted", Type: field.TypeString}, @@ -768,7 +768,7 @@ var ( {Name: "created_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "updated_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "name", Type: field.TypeString, Size: 100}, - {Name: "provider", Type: field.TypeEnum, Enums: []string{"openai", "anthropic", "gemini"}}, + {Name: "provider", Type: field.TypeEnum, Enums: []string{"openai", "anthropic", "gemini", "grok"}}, {Name: "api_mode", Type: field.TypeString, Size: 32, Default: "chat_completions"}, {Name: "description", Type: field.TypeString, Nullable: true, Size: 500, Default: ""}, {Name: "extra_headers", Type: field.TypeJSON}, @@ -865,6 +865,7 @@ var ( {Name: "video_price_480p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "video_price_720p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "video_price_1080p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "web_search_price_per_call", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "claude_code_only", Type: field.TypeBool, Default: false}, {Name: "fallback_group_id", Type: field.TypeInt64, Nullable: true}, {Name: "fallback_group_id_on_invalid_request", Type: field.TypeInt64, Nullable: true}, @@ -915,7 +916,7 @@ var ( { Name: "group_sort_order", Unique: false, - Columns: []*schema.Column{GroupsColumns[40]}, + Columns: []*schema.Column{GroupsColumns[41]}, }, }, } @@ -1559,6 +1560,7 @@ var ( {Name: "total_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}}, {Name: "actual_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}}, {Name: "rate_multiplier", Type: field.TypeFloat64, Default: 1, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "long_context_billing_applied", Type: field.TypeBool, Default: false}, {Name: "account_rate_multiplier", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, {Name: "billing_type", Type: field.TypeInt8, Default: 0}, {Name: "stream", Type: field.TypeBool, Default: false}, @@ -1591,31 +1593,31 @@ var ( ForeignKeys: []*schema.ForeignKey{ { Symbol: "usage_logs_api_keys_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[41]}, RefColumns: []*schema.Column{APIKeysColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_accounts_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[42]}, RefColumns: []*schema.Column{AccountsColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_groups_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, RefColumns: []*schema.Column{GroupsColumns[0]}, OnDelete: schema.SetNull, }, { Symbol: "usage_logs_users_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, RefColumns: []*schema.Column{UsersColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_user_subscriptions_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, RefColumns: []*schema.Column{UserSubscriptionsColumns[0]}, OnDelete: schema.SetNull, }, @@ -1624,32 +1626,32 @@ var ( { Name: "usagelog_user_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, }, { Name: "usagelog_api_key_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[41]}, }, { Name: "usagelog_account_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[42]}, }, { Name: "usagelog_group_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, }, { Name: "usagelog_subscription_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, }, { Name: "usagelog_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[40]}, }, { Name: "usagelog_model", @@ -1669,17 +1671,17 @@ var ( { Name: "usagelog_user_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[44], UsageLogsColumns[40]}, }, { Name: "usagelog_api_key_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[40], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[41], UsageLogsColumns[40]}, }, { Name: "usagelog_group_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[42], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[40]}, }, }, } diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 8d32773050..fb35531878 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -20842,6 +20842,8 @@ type GroupMutation struct { addvideo_price_720p *float64 video_price_1080p *float64 addvideo_price_1080p *float64 + web_search_price_per_call *float64 + addweb_search_price_per_call *float64 claude_code_only *bool fallback_group_id *int64 addfallback_group_id *int64 @@ -22608,6 +22610,76 @@ func (m *GroupMutation) ResetVideoPrice1080p() { delete(m.clearedFields, group.FieldVideoPrice1080p) } +// SetWebSearchPricePerCall sets the "web_search_price_per_call" field. +func (m *GroupMutation) SetWebSearchPricePerCall(f float64) { + m.web_search_price_per_call = &f + m.addweb_search_price_per_call = nil +} + +// WebSearchPricePerCall returns the value of the "web_search_price_per_call" field in the mutation. +func (m *GroupMutation) WebSearchPricePerCall() (r float64, exists bool) { + v := m.web_search_price_per_call + if v == nil { + return + } + return *v, true +} + +// OldWebSearchPricePerCall returns the old "web_search_price_per_call" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldWebSearchPricePerCall(ctx context.Context) (v *float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldWebSearchPricePerCall is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldWebSearchPricePerCall requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldWebSearchPricePerCall: %w", err) + } + return oldValue.WebSearchPricePerCall, nil +} + +// AddWebSearchPricePerCall adds f to the "web_search_price_per_call" field. +func (m *GroupMutation) AddWebSearchPricePerCall(f float64) { + if m.addweb_search_price_per_call != nil { + *m.addweb_search_price_per_call += f + } else { + m.addweb_search_price_per_call = &f + } +} + +// AddedWebSearchPricePerCall returns the value that was added to the "web_search_price_per_call" field in this mutation. +func (m *GroupMutation) AddedWebSearchPricePerCall() (r float64, exists bool) { + v := m.addweb_search_price_per_call + if v == nil { + return + } + return *v, true +} + +// ClearWebSearchPricePerCall clears the value of the "web_search_price_per_call" field. +func (m *GroupMutation) ClearWebSearchPricePerCall() { + m.web_search_price_per_call = nil + m.addweb_search_price_per_call = nil + m.clearedFields[group.FieldWebSearchPricePerCall] = struct{}{} +} + +// WebSearchPricePerCallCleared returns if the "web_search_price_per_call" field was cleared in this mutation. +func (m *GroupMutation) WebSearchPricePerCallCleared() bool { + _, ok := m.clearedFields[group.FieldWebSearchPricePerCall] + return ok +} + +// ResetWebSearchPricePerCall resets all changes to the "web_search_price_per_call" field. +func (m *GroupMutation) ResetWebSearchPricePerCall() { + m.web_search_price_per_call = nil + m.addweb_search_price_per_call = nil + delete(m.clearedFields, group.FieldWebSearchPricePerCall) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (m *GroupMutation) SetClaudeCodeOnly(b bool) { m.claude_code_only = &b @@ -23642,7 +23714,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 47) + fields := make([]string, 0, 48) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -23739,6 +23811,9 @@ func (m *GroupMutation) Fields() []string { if m.video_price_1080p != nil { fields = append(fields, group.FieldVideoPrice1080p) } + if m.web_search_price_per_call != nil { + fields = append(fields, group.FieldWebSearchPricePerCall) + } if m.claude_code_only != nil { fields = append(fields, group.FieldClaudeCodeOnly) } @@ -23856,6 +23931,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.VideoPrice720p() case group.FieldVideoPrice1080p: return m.VideoPrice1080p() + case group.FieldWebSearchPricePerCall: + return m.WebSearchPricePerCall() case group.FieldClaudeCodeOnly: return m.ClaudeCodeOnly() case group.FieldFallbackGroupID: @@ -23959,6 +24036,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldVideoPrice720p(ctx) case group.FieldVideoPrice1080p: return m.OldVideoPrice1080p(ctx) + case group.FieldWebSearchPricePerCall: + return m.OldWebSearchPricePerCall(ctx) case group.FieldClaudeCodeOnly: return m.OldClaudeCodeOnly(ctx) case group.FieldFallbackGroupID: @@ -24222,6 +24301,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetVideoPrice1080p(v) return nil + case group.FieldWebSearchPricePerCall: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetWebSearchPricePerCall(v) + return nil case group.FieldClaudeCodeOnly: v, ok := value.(bool) if !ok { @@ -24383,6 +24469,9 @@ func (m *GroupMutation) AddedFields() []string { if m.addvideo_price_1080p != nil { fields = append(fields, group.FieldVideoPrice1080p) } + if m.addweb_search_price_per_call != nil { + fields = append(fields, group.FieldWebSearchPricePerCall) + } if m.addfallback_group_id != nil { fields = append(fields, group.FieldFallbackGroupID) } @@ -24435,6 +24524,8 @@ func (m *GroupMutation) AddedField(name string) (ent.Value, bool) { return m.AddedVideoPrice720p() case group.FieldVideoPrice1080p: return m.AddedVideoPrice1080p() + case group.FieldWebSearchPricePerCall: + return m.AddedWebSearchPricePerCall() case group.FieldFallbackGroupID: return m.AddedFallbackGroupID() case group.FieldFallbackGroupIDOnInvalidRequest: @@ -24564,6 +24655,13 @@ func (m *GroupMutation) AddField(name string, value ent.Value) error { } m.AddVideoPrice1080p(v) return nil + case group.FieldWebSearchPricePerCall: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddWebSearchPricePerCall(v) + return nil case group.FieldFallbackGroupID: v, ok := value.(int64) if !ok { @@ -24633,6 +24731,9 @@ func (m *GroupMutation) ClearedFields() []string { if m.FieldCleared(group.FieldVideoPrice1080p) { fields = append(fields, group.FieldVideoPrice1080p) } + if m.FieldCleared(group.FieldWebSearchPricePerCall) { + fields = append(fields, group.FieldWebSearchPricePerCall) + } if m.FieldCleared(group.FieldFallbackGroupID) { fields = append(fields, group.FieldFallbackGroupID) } @@ -24689,6 +24790,9 @@ func (m *GroupMutation) ClearField(name string) error { case group.FieldVideoPrice1080p: m.ClearVideoPrice1080p() return nil + case group.FieldWebSearchPricePerCall: + m.ClearWebSearchPricePerCall() + return nil case group.FieldFallbackGroupID: m.ClearFallbackGroupID() return nil @@ -24802,6 +24906,9 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldVideoPrice1080p: m.ResetVideoPrice1080p() return nil + case group.FieldWebSearchPricePerCall: + m.ResetWebSearchPricePerCall() + return nil case group.FieldClaudeCodeOnly: m.ResetClaudeCodeOnly() return nil @@ -41656,83 +41763,84 @@ func (m *UsageCleanupTaskMutation) ResetEdge(name string) error { // UsageLogMutation represents an operation that mutates the UsageLog nodes in the graph. type UsageLogMutation struct { config - op Op - typ string - id *int64 - request_id *string - model *string - requested_model *string - upstream_model *string - channel_id *int64 - addchannel_id *int64 - model_mapping_chain *string - billing_tier *string - billing_mode *string - input_tokens *int - addinput_tokens *int - output_tokens *int - addoutput_tokens *int - cache_creation_tokens *int - addcache_creation_tokens *int - cache_read_tokens *int - addcache_read_tokens *int - cache_creation_5m_tokens *int - addcache_creation_5m_tokens *int - cache_creation_1h_tokens *int - addcache_creation_1h_tokens *int - input_cost *float64 - addinput_cost *float64 - output_cost *float64 - addoutput_cost *float64 - cache_creation_cost *float64 - addcache_creation_cost *float64 - cache_read_cost *float64 - addcache_read_cost *float64 - total_cost *float64 - addtotal_cost *float64 - actual_cost *float64 - addactual_cost *float64 - rate_multiplier *float64 - addrate_multiplier *float64 - account_rate_multiplier *float64 - addaccount_rate_multiplier *float64 - billing_type *int8 - addbilling_type *int8 - stream *bool - duration_ms *int - addduration_ms *int - first_token_ms *int - addfirst_token_ms *int - user_agent *string - ip_address *string - image_count *int - addimage_count *int - image_size *string - image_input_size *string - image_output_size *string - image_size_source *string - image_size_breakdown *map[string]int - video_count *int - addvideo_count *int - video_resolution *string - video_duration_seconds *int - addvideo_duration_seconds *int - cache_ttl_overridden *bool - created_at *time.Time - clearedFields map[string]struct{} - user *int64 - cleareduser bool - api_key *int64 - clearedapi_key bool - account *int64 - clearedaccount bool - group *int64 - clearedgroup bool - subscription *int64 - clearedsubscription bool - done bool - oldValue func(context.Context) (*UsageLog, error) - predicates []predicate.UsageLog + op Op + typ string + id *int64 + request_id *string + model *string + requested_model *string + upstream_model *string + channel_id *int64 + addchannel_id *int64 + model_mapping_chain *string + billing_tier *string + billing_mode *string + input_tokens *int + addinput_tokens *int + output_tokens *int + addoutput_tokens *int + cache_creation_tokens *int + addcache_creation_tokens *int + cache_read_tokens *int + addcache_read_tokens *int + cache_creation_5m_tokens *int + addcache_creation_5m_tokens *int + cache_creation_1h_tokens *int + addcache_creation_1h_tokens *int + input_cost *float64 + addinput_cost *float64 + output_cost *float64 + addoutput_cost *float64 + cache_creation_cost *float64 + addcache_creation_cost *float64 + cache_read_cost *float64 + addcache_read_cost *float64 + total_cost *float64 + addtotal_cost *float64 + actual_cost *float64 + addactual_cost *float64 + rate_multiplier *float64 + addrate_multiplier *float64 + long_context_billing_applied *bool + account_rate_multiplier *float64 + addaccount_rate_multiplier *float64 + billing_type *int8 + addbilling_type *int8 + stream *bool + duration_ms *int + addduration_ms *int + first_token_ms *int + addfirst_token_ms *int + user_agent *string + ip_address *string + image_count *int + addimage_count *int + image_size *string + image_input_size *string + image_output_size *string + image_size_source *string + image_size_breakdown *map[string]int + video_count *int + addvideo_count *int + video_resolution *string + video_duration_seconds *int + addvideo_duration_seconds *int + cache_ttl_overridden *bool + created_at *time.Time + clearedFields map[string]struct{} + user *int64 + cleareduser bool + api_key *int64 + clearedapi_key bool + account *int64 + clearedaccount bool + group *int64 + clearedgroup bool + subscription *int64 + clearedsubscription bool + done bool + oldValue func(context.Context) (*UsageLog, error) + predicates []predicate.UsageLog } var _ ent.Mutation = (*UsageLogMutation)(nil) @@ -43154,6 +43262,42 @@ func (m *UsageLogMutation) ResetRateMultiplier() { m.addrate_multiplier = nil } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (m *UsageLogMutation) SetLongContextBillingApplied(b bool) { + m.long_context_billing_applied = &b +} + +// LongContextBillingApplied returns the value of the "long_context_billing_applied" field in the mutation. +func (m *UsageLogMutation) LongContextBillingApplied() (r bool, exists bool) { + v := m.long_context_billing_applied + if v == nil { + return + } + return *v, true +} + +// OldLongContextBillingApplied returns the old "long_context_billing_applied" field's value of the UsageLog entity. +// If the UsageLog object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *UsageLogMutation) OldLongContextBillingApplied(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldLongContextBillingApplied is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldLongContextBillingApplied requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldLongContextBillingApplied: %w", err) + } + return oldValue.LongContextBillingApplied, nil +} + +// ResetLongContextBillingApplied resets all changes to the "long_context_billing_applied" field. +func (m *UsageLogMutation) ResetLongContextBillingApplied() { + m.long_context_billing_applied = nil +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (m *UsageLogMutation) SetAccountRateMultiplier(f float64) { m.account_rate_multiplier = &f @@ -44271,7 +44415,7 @@ func (m *UsageLogMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *UsageLogMutation) Fields() []string { - fields := make([]string, 0, 44) + fields := make([]string, 0, 45) if m.user != nil { fields = append(fields, usagelog.FieldUserID) } @@ -44350,6 +44494,9 @@ func (m *UsageLogMutation) Fields() []string { if m.rate_multiplier != nil { fields = append(fields, usagelog.FieldRateMultiplier) } + if m.long_context_billing_applied != nil { + fields = append(fields, usagelog.FieldLongContextBillingApplied) + } if m.account_rate_multiplier != nil { fields = append(fields, usagelog.FieldAccountRateMultiplier) } @@ -44464,6 +44611,8 @@ func (m *UsageLogMutation) Field(name string) (ent.Value, bool) { return m.ActualCost() case usagelog.FieldRateMultiplier: return m.RateMultiplier() + case usagelog.FieldLongContextBillingApplied: + return m.LongContextBillingApplied() case usagelog.FieldAccountRateMultiplier: return m.AccountRateMultiplier() case usagelog.FieldBillingType: @@ -44561,6 +44710,8 @@ func (m *UsageLogMutation) OldField(ctx context.Context, name string) (ent.Value return m.OldActualCost(ctx) case usagelog.FieldRateMultiplier: return m.OldRateMultiplier(ctx) + case usagelog.FieldLongContextBillingApplied: + return m.OldLongContextBillingApplied(ctx) case usagelog.FieldAccountRateMultiplier: return m.OldAccountRateMultiplier(ctx) case usagelog.FieldBillingType: @@ -44788,6 +44939,13 @@ func (m *UsageLogMutation) SetField(name string, value ent.Value) error { } m.SetRateMultiplier(v) return nil + case usagelog.FieldLongContextBillingApplied: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetLongContextBillingApplied(v) + return nil case usagelog.FieldAccountRateMultiplier: v, ok := value.(float64) if !ok { @@ -45419,6 +45577,9 @@ func (m *UsageLogMutation) ResetField(name string) error { case usagelog.FieldRateMultiplier: m.ResetRateMultiplier() return nil + case usagelog.FieldLongContextBillingApplied: + m.ResetLongContextBillingApplied() + return nil case usagelog.FieldAccountRateMultiplier: m.ResetAccountRateMultiplier() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index d47e7d143b..867f1cbdbd 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1044,53 +1044,53 @@ func init() { // group.DefaultVideoRateMultiplier holds the default value on creation for the video_rate_multiplier field. group.DefaultVideoRateMultiplier = groupDescVideoRateMultiplier.Default.(float64) // groupDescClaudeCodeOnly is the schema descriptor for claude_code_only field. - groupDescClaudeCodeOnly := groupFields[29].Descriptor() + groupDescClaudeCodeOnly := groupFields[30].Descriptor() // group.DefaultClaudeCodeOnly holds the default value on creation for the claude_code_only field. group.DefaultClaudeCodeOnly = groupDescClaudeCodeOnly.Default.(bool) // groupDescModelRoutingEnabled is the schema descriptor for model_routing_enabled field. - groupDescModelRoutingEnabled := groupFields[33].Descriptor() + groupDescModelRoutingEnabled := groupFields[34].Descriptor() // group.DefaultModelRoutingEnabled holds the default value on creation for the model_routing_enabled field. group.DefaultModelRoutingEnabled = groupDescModelRoutingEnabled.Default.(bool) // groupDescMcpXMLInject is the schema descriptor for mcp_xml_inject field. - groupDescMcpXMLInject := groupFields[34].Descriptor() + groupDescMcpXMLInject := groupFields[35].Descriptor() // group.DefaultMcpXMLInject holds the default value on creation for the mcp_xml_inject field. group.DefaultMcpXMLInject = groupDescMcpXMLInject.Default.(bool) // groupDescSupportedModelScopes is the schema descriptor for supported_model_scopes field. - groupDescSupportedModelScopes := groupFields[35].Descriptor() + groupDescSupportedModelScopes := groupFields[36].Descriptor() // group.DefaultSupportedModelScopes holds the default value on creation for the supported_model_scopes field. group.DefaultSupportedModelScopes = groupDescSupportedModelScopes.Default.([]string) // groupDescSortOrder is the schema descriptor for sort_order field. - groupDescSortOrder := groupFields[36].Descriptor() + groupDescSortOrder := groupFields[37].Descriptor() // group.DefaultSortOrder holds the default value on creation for the sort_order field. group.DefaultSortOrder = groupDescSortOrder.Default.(int) // groupDescAllowMessagesDispatch is the schema descriptor for allow_messages_dispatch field. - groupDescAllowMessagesDispatch := groupFields[37].Descriptor() + groupDescAllowMessagesDispatch := groupFields[38].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. - groupDescRequireOauthOnly := groupFields[38].Descriptor() + groupDescRequireOauthOnly := groupFields[39].Descriptor() // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. - groupDescRequirePrivacySet := groupFields[39].Descriptor() + groupDescRequirePrivacySet := groupFields[40].Descriptor() // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[40].Descriptor() + groupDescDefaultMappedModel := groupFields[41].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. - groupDescMessagesDispatchModelConfig := groupFields[41].Descriptor() + groupDescMessagesDispatchModelConfig := groupFields[42].Descriptor() // group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field. group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig) // groupDescModelsListConfig is the schema descriptor for models_list_config field. - groupDescModelsListConfig := groupFields[42].Descriptor() + groupDescModelsListConfig := groupFields[43].Descriptor() // group.DefaultModelsListConfig holds the default value on creation for the models_list_config field. group.DefaultModelsListConfig = groupDescModelsListConfig.Default.(domain.GroupModelsListConfig) // groupDescRpmLimit is the schema descriptor for rpm_limit field. - groupDescRpmLimit := groupFields[43].Descriptor() + groupDescRpmLimit := groupFields[44].Descriptor() // group.DefaultRpmLimit holds the default value on creation for the rpm_limit field. group.DefaultRpmLimit = groupDescRpmLimit.Default.(int) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() @@ -1940,56 +1940,60 @@ func init() { usagelogDescRateMultiplier := usagelogFields[25].Descriptor() // usagelog.DefaultRateMultiplier holds the default value on creation for the rate_multiplier field. usagelog.DefaultRateMultiplier = usagelogDescRateMultiplier.Default.(float64) + // usagelogDescLongContextBillingApplied is the schema descriptor for long_context_billing_applied field. + usagelogDescLongContextBillingApplied := usagelogFields[26].Descriptor() + // usagelog.DefaultLongContextBillingApplied holds the default value on creation for the long_context_billing_applied field. + usagelog.DefaultLongContextBillingApplied = usagelogDescLongContextBillingApplied.Default.(bool) // usagelogDescBillingType is the schema descriptor for billing_type field. - usagelogDescBillingType := usagelogFields[27].Descriptor() + usagelogDescBillingType := usagelogFields[28].Descriptor() // usagelog.DefaultBillingType holds the default value on creation for the billing_type field. usagelog.DefaultBillingType = usagelogDescBillingType.Default.(int8) // usagelogDescStream is the schema descriptor for stream field. - usagelogDescStream := usagelogFields[28].Descriptor() + usagelogDescStream := usagelogFields[29].Descriptor() // usagelog.DefaultStream holds the default value on creation for the stream field. usagelog.DefaultStream = usagelogDescStream.Default.(bool) // usagelogDescUserAgent is the schema descriptor for user_agent field. - usagelogDescUserAgent := usagelogFields[31].Descriptor() + usagelogDescUserAgent := usagelogFields[32].Descriptor() // usagelog.UserAgentValidator is a validator for the "user_agent" field. It is called by the builders before save. usagelog.UserAgentValidator = usagelogDescUserAgent.Validators[0].(func(string) error) // usagelogDescIPAddress is the schema descriptor for ip_address field. - usagelogDescIPAddress := usagelogFields[32].Descriptor() + usagelogDescIPAddress := usagelogFields[33].Descriptor() // usagelog.IPAddressValidator is a validator for the "ip_address" field. It is called by the builders before save. usagelog.IPAddressValidator = usagelogDescIPAddress.Validators[0].(func(string) error) // usagelogDescImageCount is the schema descriptor for image_count field. - usagelogDescImageCount := usagelogFields[33].Descriptor() + usagelogDescImageCount := usagelogFields[34].Descriptor() // usagelog.DefaultImageCount holds the default value on creation for the image_count field. usagelog.DefaultImageCount = usagelogDescImageCount.Default.(int) // usagelogDescImageSize is the schema descriptor for image_size field. - usagelogDescImageSize := usagelogFields[34].Descriptor() + usagelogDescImageSize := usagelogFields[35].Descriptor() // usagelog.ImageSizeValidator is a validator for the "image_size" field. It is called by the builders before save. usagelog.ImageSizeValidator = usagelogDescImageSize.Validators[0].(func(string) error) // usagelogDescImageInputSize is the schema descriptor for image_input_size field. - usagelogDescImageInputSize := usagelogFields[35].Descriptor() + usagelogDescImageInputSize := usagelogFields[36].Descriptor() // usagelog.ImageInputSizeValidator is a validator for the "image_input_size" field. It is called by the builders before save. usagelog.ImageInputSizeValidator = usagelogDescImageInputSize.Validators[0].(func(string) error) // usagelogDescImageOutputSize is the schema descriptor for image_output_size field. - usagelogDescImageOutputSize := usagelogFields[36].Descriptor() + usagelogDescImageOutputSize := usagelogFields[37].Descriptor() // usagelog.ImageOutputSizeValidator is a validator for the "image_output_size" field. It is called by the builders before save. usagelog.ImageOutputSizeValidator = usagelogDescImageOutputSize.Validators[0].(func(string) error) // usagelogDescImageSizeSource is the schema descriptor for image_size_source field. - usagelogDescImageSizeSource := usagelogFields[37].Descriptor() + usagelogDescImageSizeSource := usagelogFields[38].Descriptor() // usagelog.ImageSizeSourceValidator is a validator for the "image_size_source" field. It is called by the builders before save. usagelog.ImageSizeSourceValidator = usagelogDescImageSizeSource.Validators[0].(func(string) error) // usagelogDescVideoCount is the schema descriptor for video_count field. - usagelogDescVideoCount := usagelogFields[39].Descriptor() + usagelogDescVideoCount := usagelogFields[40].Descriptor() // usagelog.DefaultVideoCount holds the default value on creation for the video_count field. usagelog.DefaultVideoCount = usagelogDescVideoCount.Default.(int) // usagelogDescVideoResolution is the schema descriptor for video_resolution field. - usagelogDescVideoResolution := usagelogFields[40].Descriptor() + usagelogDescVideoResolution := usagelogFields[41].Descriptor() // usagelog.VideoResolutionValidator is a validator for the "video_resolution" field. It is called by the builders before save. usagelog.VideoResolutionValidator = usagelogDescVideoResolution.Validators[0].(func(string) error) // usagelogDescCacheTTLOverridden is the schema descriptor for cache_ttl_overridden field. - usagelogDescCacheTTLOverridden := usagelogFields[42].Descriptor() + usagelogDescCacheTTLOverridden := usagelogFields[43].Descriptor() // usagelog.DefaultCacheTTLOverridden holds the default value on creation for the cache_ttl_overridden field. usagelog.DefaultCacheTTLOverridden = usagelogDescCacheTTLOverridden.Default.(bool) // usagelogDescCreatedAt is the schema descriptor for created_at field. - usagelogDescCreatedAt := usagelogFields[43].Descriptor() + usagelogDescCreatedAt := usagelogFields[44].Descriptor() // usagelog.DefaultCreatedAt holds the default value on creation for the created_at field. usagelog.DefaultCreatedAt = usagelogDescCreatedAt.Default.(func() time.Time) userMixin := schema.User{}.Mixin() diff --git a/backend/ent/schema/channel_monitor.go b/backend/ent/schema/channel_monitor.go index d9594ab39c..cb62079316 100644 --- a/backend/ent/schema/channel_monitor.go +++ b/backend/ent/schema/channel_monitor.go @@ -35,7 +35,7 @@ func (ChannelMonitor) Fields() []ent.Field { NotEmpty(). MaxLen(100), field.Enum("provider"). - Values("openai", "anthropic", "gemini"), + Values("openai", "anthropic", "gemini", "grok"), field.String("api_mode"). Default("chat_completions"). MaxLen(32). diff --git a/backend/ent/schema/channel_monitor_request_template.go b/backend/ent/schema/channel_monitor_request_template.go index 0e0ce3a0b5..cf7fe05158 100644 --- a/backend/ent/schema/channel_monitor_request_template.go +++ b/backend/ent/schema/channel_monitor_request_template.go @@ -39,7 +39,7 @@ func (ChannelMonitorRequestTemplate) Fields() []ent.Field { NotEmpty(). MaxLen(100), field.Enum("provider"). - Values("openai", "anthropic", "gemini"), + Values("openai", "anthropic", "gemini", "grok"), field.String("api_mode"). Default("chat_completions"). MaxLen(32). diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index b104609a1b..70093a3cfa 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -142,6 +142,11 @@ func (Group) Fields() []ent.Field { Optional(). Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}), + field.Float("web_search_price_per_call"). + Optional(). + Nillable(). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Comment("Codex alpha/search 网页搜索单次价格(USD/次);nil 表示使用默认价 0.01(官方 $10/1000 次)"), // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). diff --git a/backend/ent/schema/usage_log.go b/backend/ent/schema/usage_log.go index e84cc1c140..6d8c2d4191 100644 --- a/backend/ent/schema/usage_log.go +++ b/backend/ent/schema/usage_log.go @@ -100,6 +100,9 @@ func (UsageLog) Fields() []ent.Field { field.Float("rate_multiplier"). Default(1). SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}), + field.Bool("long_context_billing_applied"). + Default(false). + Comment("Whether long-context pricing changed token prices for this request"), // account_rate_multiplier: 账号计费倍率快照(NULL 表示按 1.0 处理) field.Float("account_rate_multiplier"). diff --git a/backend/ent/usagelog.go b/backend/ent/usagelog.go index 4d374a8495..b13e29b2f7 100644 --- a/backend/ent/usagelog.go +++ b/backend/ent/usagelog.go @@ -75,6 +75,8 @@ type UsageLog struct { ActualCost float64 `json:"actual_cost,omitempty"` // RateMultiplier holds the value of the "rate_multiplier" field. RateMultiplier float64 `json:"rate_multiplier,omitempty"` + // Whether long-context pricing changed token prices for this request + LongContextBillingApplied bool `json:"long_context_billing_applied,omitempty"` // AccountRateMultiplier holds the value of the "account_rate_multiplier" field. AccountRateMultiplier *float64 `json:"account_rate_multiplier,omitempty"` // BillingType holds the value of the "billing_type" field. @@ -196,7 +198,7 @@ func (*UsageLog) scanValues(columns []string) ([]any, error) { switch columns[i] { case usagelog.FieldImageSizeBreakdown: values[i] = new([]byte) - case usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: + case usagelog.FieldLongContextBillingApplied, usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: values[i] = new(sql.NullBool) case usagelog.FieldInputCost, usagelog.FieldOutputCost, usagelog.FieldCacheCreationCost, usagelog.FieldCacheReadCost, usagelog.FieldTotalCost, usagelog.FieldActualCost, usagelog.FieldRateMultiplier, usagelog.FieldAccountRateMultiplier: values[i] = new(sql.NullFloat64) @@ -391,6 +393,12 @@ func (_m *UsageLog) assignValues(columns []string, values []any) error { } else if value.Valid { _m.RateMultiplier = value.Float64 } + case usagelog.FieldLongContextBillingApplied: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field long_context_billing_applied", values[i]) + } else if value.Valid { + _m.LongContextBillingApplied = value.Bool + } case usagelog.FieldAccountRateMultiplier: if value, ok := values[i].(*sql.NullFloat64); !ok { return fmt.Errorf("unexpected type %T for field account_rate_multiplier", values[i]) @@ -667,6 +675,9 @@ func (_m *UsageLog) String() string { builder.WriteString("rate_multiplier=") builder.WriteString(fmt.Sprintf("%v", _m.RateMultiplier)) builder.WriteString(", ") + builder.WriteString("long_context_billing_applied=") + builder.WriteString(fmt.Sprintf("%v", _m.LongContextBillingApplied)) + builder.WriteString(", ") if v := _m.AccountRateMultiplier; v != nil { builder.WriteString("account_rate_multiplier=") builder.WriteString(fmt.Sprintf("%v", *v)) diff --git a/backend/ent/usagelog/usagelog.go b/backend/ent/usagelog/usagelog.go index a74a92c40f..a87d937195 100644 --- a/backend/ent/usagelog/usagelog.go +++ b/backend/ent/usagelog/usagelog.go @@ -66,6 +66,8 @@ const ( FieldActualCost = "actual_cost" // FieldRateMultiplier holds the string denoting the rate_multiplier field in the database. FieldRateMultiplier = "rate_multiplier" + // FieldLongContextBillingApplied holds the string denoting the long_context_billing_applied field in the database. + FieldLongContextBillingApplied = "long_context_billing_applied" // FieldAccountRateMultiplier holds the string denoting the account_rate_multiplier field in the database. FieldAccountRateMultiplier = "account_rate_multiplier" // FieldBillingType holds the string denoting the billing_type field in the database. @@ -180,6 +182,7 @@ var Columns = []string{ FieldTotalCost, FieldActualCost, FieldRateMultiplier, + FieldLongContextBillingApplied, FieldAccountRateMultiplier, FieldBillingType, FieldStream, @@ -251,6 +254,8 @@ var ( DefaultActualCost float64 // DefaultRateMultiplier holds the default value on creation for the "rate_multiplier" field. DefaultRateMultiplier float64 + // DefaultLongContextBillingApplied holds the default value on creation for the "long_context_billing_applied" field. + DefaultLongContextBillingApplied bool // DefaultBillingType holds the default value on creation for the "billing_type" field. DefaultBillingType int8 // DefaultStream holds the default value on creation for the "stream" field. @@ -417,6 +422,11 @@ func ByRateMultiplier(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldRateMultiplier, opts...).ToFunc() } +// ByLongContextBillingApplied orders the results by the long_context_billing_applied field. +func ByLongContextBillingApplied(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldLongContextBillingApplied, opts...).ToFunc() +} + // ByAccountRateMultiplier orders the results by the account_rate_multiplier field. func ByAccountRateMultiplier(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAccountRateMultiplier, opts...).ToFunc() diff --git a/backend/ent/usagelog/where.go b/backend/ent/usagelog/where.go index 4b08cc3425..a9462e0d0e 100644 --- a/backend/ent/usagelog/where.go +++ b/backend/ent/usagelog/where.go @@ -185,6 +185,11 @@ func RateMultiplier(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldRateMultiplier, v)) } +// LongContextBillingApplied applies equality check predicate on the "long_context_billing_applied" field. It's identical to LongContextBillingAppliedEQ. +func LongContextBillingApplied(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v)) +} + // AccountRateMultiplier applies equality check predicate on the "account_rate_multiplier" field. It's identical to AccountRateMultiplierEQ. func AccountRateMultiplier(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v)) @@ -1465,6 +1470,16 @@ func RateMultiplierLTE(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldLTE(FieldRateMultiplier, v)) } +// LongContextBillingAppliedEQ applies the EQ predicate on the "long_context_billing_applied" field. +func LongContextBillingAppliedEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v)) +} + +// LongContextBillingAppliedNEQ applies the NEQ predicate on the "long_context_billing_applied" field. +func LongContextBillingAppliedNEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldNEQ(FieldLongContextBillingApplied, v)) +} + // AccountRateMultiplierEQ applies the EQ predicate on the "account_rate_multiplier" field. func AccountRateMultiplierEQ(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v)) diff --git a/backend/ent/usagelog_create.go b/backend/ent/usagelog_create.go index 3326f72fc0..31cf45328e 100644 --- a/backend/ent/usagelog_create.go +++ b/backend/ent/usagelog_create.go @@ -351,6 +351,20 @@ func (_c *UsageLogCreate) SetNillableRateMultiplier(v *float64) *UsageLogCreate return _c } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_c *UsageLogCreate) SetLongContextBillingApplied(v bool) *UsageLogCreate { + _c.mutation.SetLongContextBillingApplied(v) + return _c +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_c *UsageLogCreate) SetNillableLongContextBillingApplied(v *bool) *UsageLogCreate { + if v != nil { + _c.SetLongContextBillingApplied(*v) + } + return _c +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_c *UsageLogCreate) SetAccountRateMultiplier(v float64) *UsageLogCreate { _c.mutation.SetAccountRateMultiplier(v) @@ -707,6 +721,10 @@ func (_c *UsageLogCreate) defaults() { v := usagelog.DefaultRateMultiplier _c.mutation.SetRateMultiplier(v) } + if _, ok := _c.mutation.LongContextBillingApplied(); !ok { + v := usagelog.DefaultLongContextBillingApplied + _c.mutation.SetLongContextBillingApplied(v) + } if _, ok := _c.mutation.BillingType(); !ok { v := usagelog.DefaultBillingType _c.mutation.SetBillingType(v) @@ -824,6 +842,9 @@ func (_c *UsageLogCreate) check() error { if _, ok := _c.mutation.RateMultiplier(); !ok { return &ValidationError{Name: "rate_multiplier", err: errors.New(`ent: missing required field "UsageLog.rate_multiplier"`)} } + if _, ok := _c.mutation.LongContextBillingApplied(); !ok { + return &ValidationError{Name: "long_context_billing_applied", err: errors.New(`ent: missing required field "UsageLog.long_context_billing_applied"`)} + } if _, ok := _c.mutation.BillingType(); !ok { return &ValidationError{Name: "billing_type", err: errors.New(`ent: missing required field "UsageLog.billing_type"`)} } @@ -997,6 +1018,10 @@ func (_c *UsageLogCreate) createSpec() (*UsageLog, *sqlgraph.CreateSpec) { _spec.SetField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) _node.RateMultiplier = value } + if value, ok := _c.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + _node.LongContextBillingApplied = value + } if value, ok := _c.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) _node.AccountRateMultiplier = &value @@ -1650,6 +1675,18 @@ func (u *UsageLogUpsert) AddRateMultiplier(v float64) *UsageLogUpsert { return u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsert) SetLongContextBillingApplied(v bool) *UsageLogUpsert { + u.Set(usagelog.FieldLongContextBillingApplied, v) + return u +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsert) UpdateLongContextBillingApplied() *UsageLogUpsert { + u.SetExcluded(usagelog.FieldLongContextBillingApplied) + return u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsert) SetAccountRateMultiplier(v float64) *UsageLogUpsert { u.Set(usagelog.FieldAccountRateMultiplier, v) @@ -2531,6 +2568,20 @@ func (u *UsageLogUpsertOne) UpdateRateMultiplier() *UsageLogUpsertOne { }) } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsertOne) SetLongContextBillingApplied(v bool) *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.SetLongContextBillingApplied(v) + }) +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsertOne) UpdateLongContextBillingApplied() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateLongContextBillingApplied() + }) +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsertOne) SetAccountRateMultiplier(v float64) *UsageLogUpsertOne { return u.Update(func(s *UsageLogUpsert) { @@ -3631,6 +3682,20 @@ func (u *UsageLogUpsertBulk) UpdateRateMultiplier() *UsageLogUpsertBulk { }) } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsertBulk) SetLongContextBillingApplied(v bool) *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.SetLongContextBillingApplied(v) + }) +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsertBulk) UpdateLongContextBillingApplied() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateLongContextBillingApplied() + }) +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsertBulk) SetAccountRateMultiplier(v float64) *UsageLogUpsertBulk { return u.Update(func(s *UsageLogUpsert) { diff --git a/backend/ent/usagelog_update.go b/backend/ent/usagelog_update.go index 00a65ccff1..2a60d6f44d 100644 --- a/backend/ent/usagelog_update.go +++ b/backend/ent/usagelog_update.go @@ -542,6 +542,20 @@ func (_u *UsageLogUpdate) AddRateMultiplier(v float64) *UsageLogUpdate { return _u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_u *UsageLogUpdate) SetLongContextBillingApplied(v bool) *UsageLogUpdate { + _u.mutation.SetLongContextBillingApplied(v) + return _u +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_u *UsageLogUpdate) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdate { + if v != nil { + _u.SetLongContextBillingApplied(*v) + } + return _u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_u *UsageLogUpdate) SetAccountRateMultiplier(v float64) *UsageLogUpdate { _u.mutation.ResetAccountRateMultiplier() @@ -1199,6 +1213,9 @@ func (_u *UsageLogUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AddedRateMultiplier(); ok { _spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) } + if value, ok := _u.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + } if value, ok := _u.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) } @@ -1982,6 +1999,20 @@ func (_u *UsageLogUpdateOne) AddRateMultiplier(v float64) *UsageLogUpdateOne { return _u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_u *UsageLogUpdateOne) SetLongContextBillingApplied(v bool) *UsageLogUpdateOne { + _u.mutation.SetLongContextBillingApplied(v) + return _u +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_u *UsageLogUpdateOne) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdateOne { + if v != nil { + _u.SetLongContextBillingApplied(*v) + } + return _u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_u *UsageLogUpdateOne) SetAccountRateMultiplier(v float64) *UsageLogUpdateOne { _u.mutation.ResetAccountRateMultiplier() @@ -2669,6 +2700,9 @@ func (_u *UsageLogUpdateOne) sqlSave(ctx context.Context) (_node *UsageLog, err if value, ok := _u.mutation.AddedRateMultiplier(); ok { _spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) } + if value, ok := _u.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + } if value, ok := _u.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) } diff --git a/backend/go.mod b/backend/go.mod index a06f06437d..64b5c30d96 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -44,6 +44,7 @@ require ( go.uber.org/zap v1.24.0 golang.org/x/crypto v0.51.0 golang.org/x/image v0.39.0 + golang.org/x/mod v0.35.0 golang.org/x/net v0.55.0 golang.org/x/sync v0.20.0 golang.org/x/term v0.43.0 @@ -176,7 +177,6 @@ require ( go.uber.org/multierr v1.9.0 // indirect golang.org/x/arch v0.3.0 // indirect golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect - golang.org/x/mod v0.35.0 // indirect golang.org/x/sys v0.45.0 // indirect golang.org/x/text v0.37.0 // indirect golang.org/x/tools v0.44.0 // indirect diff --git a/backend/go.sum b/backend/go.sum index 4738443bb9..3d9989bb43 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -220,6 +220,8 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-runewidth v0.0.15 h1:UNAjwbU9l54TA3KzvqLGxwWjHmMgBUVhBiTjelZgg3U= +github.com/mattn/go-runewidth v0.0.15/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM= github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= github.com/mdelapenya/tlscert v0.2.0 h1:7H81W6Z/4weDvZBNOfQte5GpIMo0lGYEeWbkGp5LJHI= @@ -253,6 +255,8 @@ github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A= github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/olekukonko/tablewriter v0.0.5 h1:P2Ga83D34wi1o9J6Wh1mRuqd4mF/x/lgBS7N7AbDhec= +github.com/olekukonko/tablewriter v0.0.5/go.mod h1:hPp6KlRPjbx+hW8ykQs1w3UBbZlj6HuIJcUGPhkA7kY= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040= @@ -282,6 +286,8 @@ github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEv github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= +github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII= @@ -314,6 +320,8 @@ github.com/spf13/afero v1.11.0 h1:WJQKhtpdm3v2IzqG8VMqrr6Rf3UYpEF239Jy9wNepM8= github.com/spf13/afero v1.11.0/go.mod h1:GH9Y3pIexgf1MTIWtNGyogA5MwRIDXGUr+hbWNoBjkY= github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0= github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo= +github.com/spf13/cobra v1.7.0 h1:hyqWnYt1ZQShIddO5kBpj3vu05/++x6tJ6dg8EC572I= +github.com/spf13/cobra v1.7.0/go.mod h1:uLxZILRyS/50WlhOIKD7W6V5bgeIt+4sICxh6uRMrb0= github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ= diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index df3afb6c7e..1262845cea 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -601,6 +601,7 @@ type ServerConfig struct { Host string `mapstructure:"host"` Port int `mapstructure:"port"` Mode string `mapstructure:"mode"` // debug/release + EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接 ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒) IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒) @@ -818,6 +819,8 @@ type GatewayConfig struct { ImageStreamDataIntervalTimeout int `mapstructure:"image_stream_data_interval_timeout"` // ImageStreamKeepaliveInterval: 图片流式 keepalive 间隔(秒),0表示禁用 ImageStreamKeepaliveInterval int `mapstructure:"image_stream_keepalive_interval"` + // ImageNonstreamKeepaliveInterval: 图片非流式 JSON keepalive 间隔(秒),0表示禁用 + ImageNonstreamKeepaliveInterval int `mapstructure:"image_nonstream_keepalive_interval"` // MaxLineSize: 上游 SSE 单行最大字节数(0使用默认值) MaxLineSize int `mapstructure:"max_line_size"` @@ -923,6 +926,12 @@ type GatewayOpenAIWSConfig struct { ModeRouterV2Enabled bool `mapstructure:"mode_router_v2_enabled"` // IngressModeDefault: ingress 默认模式(off/ctx_pool/passthrough/http_bridge) IngressModeDefault string `mapstructure:"ingress_mode_default"` + // IngressInterTurnIdleTimeoutSeconds bounds the time a client may remain idle + // between completed ingress turns. Zero disables this protection. + IngressInterTurnIdleTimeoutSeconds int `mapstructure:"ingress_inter_turn_idle_timeout_seconds"` + // MaxIngressConnectionsPerAPIKey bounds live client WebSocket ingress sessions + // per API key across all instances. Zero disables this protection. + MaxIngressConnectionsPerAPIKey int `mapstructure:"max_ingress_connections_per_api_key"` // Enabled: 全局总开关(默认 true) Enabled bool `mapstructure:"enabled"` // OAuthEnabled: 是否允许 OpenAI OAuth 账号使用 WS @@ -1453,6 +1462,9 @@ func load(allowMissingJWTSecret bool) (*Config, error) { // 环境变量支持 viper.AutomaticEnv() viper.SetEnvKeyReplacer(strings.NewReplacer(".", "_")) + if err := viper.BindEnv("server.enable_server_timing", "ENABLE_SERVER_TIMING"); err != nil { + return nil, fmt.Errorf("bind ENABLE_SERVER_TIMING: %w", err) + } // 默认值 setDefaults() @@ -1608,6 +1620,7 @@ func setDefaults() { viper.SetDefault("server.host", "0.0.0.0") viper.SetDefault("server.port", 8080) viper.SetDefault("server.mode", "release") + viper.SetDefault("server.enable_server_timing", false) viper.SetDefault("server.frontend_url", "") viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头 viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时 @@ -1945,6 +1958,8 @@ func setDefaults() { viper.SetDefault("gateway.openai_ws.enabled", true) viper.SetDefault("gateway.openai_ws.mode_router_v2_enabled", false) viper.SetDefault("gateway.openai_ws.ingress_mode_default", "ctx_pool") + viper.SetDefault("gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds", 300) + viper.SetDefault("gateway.openai_ws.max_ingress_connections_per_api_key", 64) viper.SetDefault("gateway.openai_ws.oauth_enabled", true) viper.SetDefault("gateway.openai_ws.apikey_enabled", true) viper.SetDefault("gateway.openai_ws.force_http", false) @@ -2024,6 +2039,7 @@ func setDefaults() { viper.SetDefault("gateway.stream_keepalive_interval", 10) viper.SetDefault("gateway.image_stream_data_interval_timeout", 900) viper.SetDefault("gateway.image_stream_keepalive_interval", 10) + viper.SetDefault("gateway.image_nonstream_keepalive_interval", 0) viper.SetDefault("gateway.max_line_size", 500*1024*1024) viper.SetDefault("gateway.scheduling.sticky_session_max_waiting", 3) viper.SetDefault("gateway.scheduling.sticky_session_wait_timeout", 120*time.Second) @@ -2710,6 +2726,13 @@ func (c *Config) Validate() error { (c.Gateway.ImageStreamKeepaliveInterval < 5 || c.Gateway.ImageStreamKeepaliveInterval > 60) { return fmt.Errorf("gateway.image_stream_keepalive_interval must be 0 or between 5-60 seconds") } + if c.Gateway.ImageNonstreamKeepaliveInterval < 0 { + return fmt.Errorf("gateway.image_nonstream_keepalive_interval must be non-negative") + } + if c.Gateway.ImageNonstreamKeepaliveInterval != 0 && + (c.Gateway.ImageNonstreamKeepaliveInterval < 5 || c.Gateway.ImageNonstreamKeepaliveInterval > 60) { + return fmt.Errorf("gateway.image_nonstream_keepalive_interval must be 0 or between 5-60 seconds") + } // 兼容旧键 sticky_previous_response_ttl_seconds if c.Gateway.OpenAIWS.StickyResponseIDTTLSeconds <= 0 && c.Gateway.OpenAIWS.StickyPreviousResponseTTLSeconds > 0 { c.Gateway.OpenAIWS.StickyResponseIDTTLSeconds = c.Gateway.OpenAIWS.StickyPreviousResponseTTLSeconds @@ -2717,6 +2740,12 @@ func (c *Config) Validate() error { if c.Gateway.OpenAIWS.MaxConnsPerAccount <= 0 { return fmt.Errorf("gateway.openai_ws.max_conns_per_account must be positive") } + if c.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds < 0 { + return fmt.Errorf("gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds must be non-negative") + } + if c.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey < 0 { + return fmt.Errorf("gateway.openai_ws.max_ingress_connections_per_api_key must be non-negative") + } if c.Gateway.OpenAIWS.MinIdlePerAccount < 0 { return fmt.Errorf("gateway.openai_ws.min_idle_per_account must be non-negative") } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 32aff543af..2f9defb80e 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -17,6 +17,23 @@ func resetViperWithJWTSecret(t *testing.T) { t.Setenv("JWT_SECRET", strings.Repeat("x", 32)) } +func TestLoadServerTimingConfig(t *testing.T) { + t.Run("disabled by default", func(t *testing.T) { + resetViperWithJWTSecret(t) + cfg, err := Load() + require.NoError(t, err) + require.False(t, cfg.Server.EnableServerTiming) + }) + + t.Run("enabled by exact environment variable", func(t *testing.T) { + resetViperWithJWTSecret(t) + t.Setenv("ENABLE_SERVER_TIMING", "true") + cfg, err := Load() + require.NoError(t, err) + require.True(t, cfg.Server.EnableServerTiming) + }) +} + func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) { viper.Reset() t.Setenv("JWT_SECRET", "") @@ -182,6 +199,12 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) { if cfg.Gateway.OpenAIWS.IngressModeDefault != "ctx_pool" { t.Fatalf("Gateway.OpenAIWS.IngressModeDefault = %q, want %q", cfg.Gateway.OpenAIWS.IngressModeDefault, "ctx_pool") } + if cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds != 300 { + t.Fatalf("Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = %d, want 300", cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) + } + if cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey != 64 { + t.Fatalf("Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = %d, want 64", cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey) + } } func TestLoadDefaultOpenAICompactModel(t *testing.T) { @@ -236,6 +259,15 @@ func TestLoadOpenAIResponseHeaderTimeoutFromEnv(t *testing.T) { require.Equal(t, 1800, cfg.Gateway.OpenAIResponseHeaderTimeout) } +func TestLoadImageNonstreamKeepaliveFromEnv(t *testing.T) { + resetViperWithJWTSecret(t) + t.Setenv("GATEWAY_IMAGE_NONSTREAM_KEEPALIVE_INTERVAL", "15") + + cfg, err := Load() + require.NoError(t, err) + require.Equal(t, 15, cfg.Gateway.ImageNonstreamKeepaliveInterval) +} + func TestLoadOpenAIWSStickyTTLCompatibility(t *testing.T) { resetViperWithJWTSecret(t) t.Setenv("GATEWAY_OPENAI_WS_STICKY_RESPONSE_ID_TTL_SECONDS", "0") @@ -1406,6 +1438,16 @@ func TestValidateConfigErrors(t *testing.T) { mutate: func(c *Config) { c.Gateway.ImageStreamKeepaliveInterval = -1 }, wantErr: "gateway.image_stream_keepalive_interval must be non-negative", }, + { + name: "gateway image nonstream keepalive range", + mutate: func(c *Config) { c.Gateway.ImageNonstreamKeepaliveInterval = 4 }, + wantErr: "gateway.image_nonstream_keepalive_interval", + }, + { + name: "gateway image nonstream keepalive negative", + mutate: func(c *Config) { c.Gateway.ImageNonstreamKeepaliveInterval = -1 }, + wantErr: "gateway.image_nonstream_keepalive_interval must be non-negative", + }, { name: "gateway image stream data interval range", mutate: func(c *Config) { c.Gateway.ImageStreamDataIntervalTimeout = 30 }, @@ -1640,6 +1682,16 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) { mutate: func(c *Config) { c.Gateway.OpenAIWS.MaxConnsPerAccount = 0 }, wantErr: "gateway.openai_ws.max_conns_per_account", }, + { + name: "ingress_inter_turn_idle_timeout_seconds 不能为负数", + mutate: func(c *Config) { c.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = -1 }, + wantErr: "gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds", + }, + { + name: "max_ingress_connections_per_api_key 不能为负数", + mutate: func(c *Config) { c.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = -1 }, + wantErr: "gateway.openai_ws.max_ingress_connections_per_api_key", + }, { name: "min_idle_per_account 不能为负数", mutate: func(c *Config) { c.Gateway.OpenAIWS.MinIdlePerAccount = -1 }, @@ -1964,6 +2016,9 @@ func TestLoad_DefaultGatewayImageStreamConfig(t *testing.T) { if cfg.Gateway.ImageStreamKeepaliveInterval != 10 { t.Fatalf("image_stream_keepalive_interval = %d, want 10", cfg.Gateway.ImageStreamKeepaliveInterval) } + if cfg.Gateway.ImageNonstreamKeepaliveInterval != 0 { + t.Fatalf("image_nonstream_keepalive_interval = %d, want 0", cfg.Gateway.ImageNonstreamKeepaliveInterval) + } if cfg.Gateway.ImageConcurrency.Enabled { t.Fatalf("image_concurrency.enabled = true, want false") } diff --git a/backend/internal/handler/admin/account_codex_import.go b/backend/internal/handler/admin/account_codex_import.go index 8abd269f23..a6a07af0c1 100644 --- a/backend/internal/handler/admin/account_codex_import.go +++ b/backend/internal/handler/admin/account_codex_import.go @@ -120,6 +120,10 @@ func (h *AccountHandler) ImportCodexSession(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.Concurrency != nil && *req.Concurrency < 0 { response.BadRequest(c, "concurrency must be >= 0") return diff --git a/backend/internal/handler/admin/account_codex_import_test.go b/backend/internal/handler/admin/account_codex_import_test.go index a52463aa86..96a033d8c3 100644 --- a/backend/internal/handler/admin/account_codex_import_test.go +++ b/backend/internal/handler/admin/account_codex_import_test.go @@ -630,6 +630,7 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) "chatgpt_user_id": "user-1", "access_token": existingToken, }, + Extra: map[string]any{"openai_long_context_billing_enabled": false}, }}) handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} @@ -650,6 +651,9 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 { t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts) } + if got := svc.updatedAccounts[0].input.Extra["openai_long_context_billing_enabled"]; got != false { + t.Fatalf("openai_long_context_billing_enabled = %v, want false", got) + } } func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) { diff --git a/backend/internal/handler/admin/account_data.go b/backend/internal/handler/admin/account_data.go index bf872c4826..e44d726fd6 100644 --- a/backend/internal/handler/admin/account_data.go +++ b/backend/internal/handler/admin/account_data.go @@ -460,6 +460,7 @@ func (h *AccountHandler) importData(ctx context.Context, req DataImportRequest) if created.Platform == service.PlatformAntigravity && created.Type == service.AccountTypeOAuth { privacyAccounts = append(privacyAccounts, created) } + h.scheduleGrokImportProbe(created) result.AccountCreated++ } diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index a4b0773999..e4ed5b46b0 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -60,6 +60,7 @@ type AccountHandler struct { sessionLimitCache service.SessionLimitCache rpmCache service.RPMCache tokenCacheInvalidator service.TokenCacheInvalidator + grokImportProber grokUsageProber } // NewAccountHandler creates a new admin account handler @@ -784,6 +785,10 @@ func (h *AccountHandler) Create(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(req.Platform, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.RateMultiplier != nil && *req.RateMultiplier < 0 { response.BadRequest(c, "rate_multiplier must be >= 0") return @@ -851,6 +856,7 @@ func (h *AccountHandler) Create(c *gin.Context) { // OpenAI APIKey 账号创建后异步探测上游 /v1/responses 能力。 // 探测失败不影响账号创建响应。 h.scheduleOpenAIResponsesProbe(createdAccount) + h.scheduleGrokImportProbe(createdAccount) response.Success(c, result.Data) } @@ -1299,6 +1305,10 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) { response.ErrorFrom(c, infraerrors.BadRequest("NOT_OAUTH", "cannot apply oauth credentials to non-OAuth account")) return } + if err := service.ValidateOpenAILongContextBillingExtra(existing.Platform, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } updatedAccount, err := h.adminService.UpdateAccount(ctx, accountID, &service.UpdateAccountInput{ Type: req.Type, @@ -1592,6 +1602,12 @@ func (h *AccountHandler) BatchCreate(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + for _, item := range req.Accounts { + if err := service.ValidateOpenAILongContextBillingExtra(item.Platform, item.Extra); err != nil { + response.ErrorFrom(c, err) + return + } + } executeAdminIdempotentJSON(c, "admin.accounts.batch_create", req, service.DefaultWriteIdempotencyTTL(), func(ctx context.Context) (any, error) { success := 0 @@ -1653,6 +1669,7 @@ func (h *AccountHandler) BatchCreate(c *gin.Context) { } // OpenAI APIKey 账号异步探测 /v1/responses 能力。 h.scheduleOpenAIResponsesProbe(account) + h.scheduleGrokImportProbe(account) success++ results = append(results, gin.H{ "name": item.Name, diff --git a/backend/internal/handler/admin/account_handler_long_context_billing_test.go b/backend/internal/handler/admin/account_handler_long_context_billing_test.go new file mode 100644 index 0000000000..d50513a3e8 --- /dev/null +++ b/backend/internal/handler/admin/account_handler_long_context_billing_test.go @@ -0,0 +1,165 @@ +package admin + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestAccountAdminBoundariesRejectMalformedOpenAILongContextBillingValue(t *testing.T) { + const malformedExtra = `"extra":{"openai_long_context_billing_enabled":"true"}` + + tests := []struct { + name string + method string + path string + body string + mount func(*gin.Engine, *AccountHandler) + setup func(*stubAdminService) + }{ + { + name: "create", + method: http.MethodPost, + path: "/accounts", + body: `{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts", handler.Create) }, + }, + { + name: "update", + method: http.MethodPut, + path: "/accounts/1", + body: `{` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.PUT("/accounts/:id", handler.Update) }, + setup: func(stub *stubAdminService) { + stub.updateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid") + }, + }, + { + name: "bulk update", + method: http.MethodPost, + path: "/accounts/bulk-update", + body: `{"account_ids":[1],` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { + router.POST("/accounts/bulk-update", handler.BulkUpdate) + }, + setup: func(stub *stubAdminService) { + stub.bulkUpdateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid") + }, + }, + { + name: "batch create", + method: http.MethodPost, + path: "/accounts/batch", + body: `{"accounts":[{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}]}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts/batch", handler.BatchCreate) }, + }, + { + name: "Codex session import", + method: http.MethodPost, + path: "/accounts/import-codex-session", + body: `{"content":"token",` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { + router.POST("/accounts/import-codex-session", handler.ImportCodexSession) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + stub := newStubAdminService() + if tt.setup != nil { + tt.setup(stub) + } + handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + tt.mount(router, handler) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body)) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) + }) + } +} + +func TestAccountCreateBoundaryDoesNotApplyOpenAIValidationToOtherPlatforms(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewAccountHandler(newStubAdminService(), nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.POST("/accounts", handler.Create) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/accounts", bytes.NewBufferString( + `{"name":"account","platform":"anthropic","type":"apikey","credentials":{"api_key":"test"},"extra":{"openai_long_context_billing_enabled":"provider-owned"}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) +} + +func TestApplyOAuthCredentialsRejectsMalformedOpenAILongContextBillingBeforeMutation(t *testing.T) { + gin.SetMode(gin.TestMode) + stub := newStubAdminService() + stub.getAccountResult = &service.Account{ + ID: 1, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + } + handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.POST("/accounts/:id/apply-oauth-credentials", handler.ApplyOAuthCredentials) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/accounts/1/apply-oauth-credentials", bytes.NewBufferString( + `{"type":"oauth","credentials":{"access_token":"new-token"},"extra":{"openai_long_context_billing_enabled":"true"}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) + require.Zero(t, stub.updateAccountCalls) + require.Zero(t, stub.updateAccountExtraCalls) +} + +func TestOpenAIOAuthCodexPATBoundaryRejectsMalformedOpenAILongContextBillingValueBeforeTokenValidation(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewOpenAIOAuthHandler(nil, newStubAdminService(), nil) + router := gin.New() + router.Use(gin.Recovery()) + router.POST("/openai/create-from-codex-pat", handler.CreateAccountFromCodexPAT) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/openai/create-from-codex-pat", bytes.NewBufferString( + `{"access_token":"token","extra":{"openai_long_context_billing_enabled":1}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) +} diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go index 7a7cbb473e..5e9c4d517e 100644 --- a/backend/internal/handler/admin/admin_service_stub_test.go +++ b/backend/internal/handler/admin/admin_service_stub_test.go @@ -33,6 +33,9 @@ type stubAdminService struct { createSparkShadowErr error updateAccountErr error bulkUpdateAccountErr error + getAccountResult *service.Account + updateAccountCalls int + updateAccountExtraCalls int checkMixedErr error lastMixedCheck struct { accountID int64 @@ -388,6 +391,9 @@ func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ cont } func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) { + if s.getAccountResult != nil { + return s.getAccountResult, nil + } account := service.Account{ID: id, Name: "account", Status: service.StatusActive} return &account, nil } @@ -413,6 +419,7 @@ func (s *stubAdminService) CreateAccount(ctx context.Context, input *service.Cre } func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) { + s.updateAccountCalls++ if s.updateAccountErr != nil { return nil, s.updateAccountErr } @@ -421,6 +428,7 @@ func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *s } func (s *stubAdminService) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error { + s.updateAccountExtraCalls++ return nil } diff --git a/backend/internal/handler/admin/channel_monitor_handler.go b/backend/internal/handler/admin/channel_monitor_handler.go index 4ef774e9e7..a69b835849 100644 --- a/backend/internal/handler/admin/channel_monitor_handler.go +++ b/backend/internal/handler/admin/channel_monitor_handler.go @@ -37,11 +37,11 @@ func NewChannelMonitorHandler(monitorService *service.ChannelMonitorService) *Ch type channelMonitorCreateRequest struct { Name string `json:"name" binding:"required,max=100"` - Provider string `json:"provider" binding:"required,oneof=openai anthropic gemini"` + Provider string `json:"provider" binding:"required,oneof=openai anthropic gemini grok"` APIMode string `json:"api_mode" binding:"omitempty,oneof=chat_completions responses"` Endpoint string `json:"endpoint" binding:"required,max=500"` APIKey string `json:"api_key" binding:"required,max=2000"` - PrimaryModel string `json:"primary_model" binding:"required,max=200"` + PrimaryModel string `json:"primary_model" binding:"max=200"` ExtraModels []string `json:"extra_models"` GroupName string `json:"group_name" binding:"max=100"` Enabled *bool `json:"enabled"` @@ -55,7 +55,7 @@ type channelMonitorCreateRequest struct { type channelMonitorUpdateRequest struct { Name *string `json:"name" binding:"omitempty,max=100"` - Provider *string `json:"provider" binding:"omitempty,oneof=openai anthropic gemini"` + Provider *string `json:"provider" binding:"omitempty,oneof=openai anthropic gemini grok"` APIMode *string `json:"api_mode" binding:"omitempty,oneof=chat_completions responses"` Endpoint *string `json:"endpoint" binding:"omitempty,max=500"` APIKey *string `json:"api_key" binding:"omitempty,max=2000"` diff --git a/backend/internal/handler/admin/channel_monitor_template_handler.go b/backend/internal/handler/admin/channel_monitor_template_handler.go index c842f465c8..497e3d195b 100644 --- a/backend/internal/handler/admin/channel_monitor_template_handler.go +++ b/backend/internal/handler/admin/channel_monitor_template_handler.go @@ -26,7 +26,7 @@ func NewChannelMonitorRequestTemplateHandler(templateService *service.ChannelMon type channelMonitorTemplateCreateRequest struct { Name string `json:"name" binding:"required,max=100"` - Provider string `json:"provider" binding:"required,oneof=openai anthropic gemini"` + Provider string `json:"provider" binding:"required,oneof=openai anthropic gemini grok"` APIMode string `json:"api_mode" binding:"omitempty,oneof=chat_completions responses"` Description string `json:"description" binding:"max=500"` ExtraHeaders map[string]string `json:"extra_headers"` diff --git a/backend/internal/handler/admin/grok_import_probe.go b/backend/internal/handler/admin/grok_import_probe.go new file mode 100644 index 0000000000..f1df15bba9 --- /dev/null +++ b/backend/internal/handler/admin/grok_import_probe.go @@ -0,0 +1,203 @@ +package admin + +import ( + "context" + "log/slog" + "sync" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +const ( + grokImportProbeConcurrency = 3 + grokImportProbeTimeout = 25 * time.Second +) + +type grokUsageProber interface { + ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) +} + +type grokImportProbeTask struct { + prober grokUsageProber + accountID int64 +} + +type grokImportProbeScheduler struct { + mu sync.Mutex + queue []grokImportProbeTask + concurrency int + workers int + maxWorkers int + timeout time.Duration +} + +var defaultGrokImportProbeScheduler = newGrokImportProbeScheduler( + grokImportProbeConcurrency, + grokImportProbeTimeout, +) + +func newGrokImportProbeScheduler(concurrency int, timeout time.Duration) *grokImportProbeScheduler { + if concurrency <= 0 { + concurrency = 1 + } + if timeout <= 0 { + timeout = grokImportProbeTimeout + } + return &grokImportProbeScheduler{ + concurrency: concurrency, + timeout: timeout, + } +} + +func (s *grokImportProbeScheduler) schedule(prober grokUsageProber, account *service.Account) { + if s == nil || prober == nil || account == nil || account.ID <= 0 { + return + } + if account.Platform != service.PlatformGrok || account.Type != service.AccountTypeOAuth { + return + } + + s.mu.Lock() + s.queue = append(s.queue, grokImportProbeTask{prober: prober, accountID: account.ID}) + if s.workers < s.concurrency { + s.workers++ + if s.workers > s.maxWorkers { + s.maxWorkers = s.workers + } + go s.worker() + } + s.mu.Unlock() +} + +func (s *grokImportProbeScheduler) worker() { + for { + task, ok := s.nextTask() + if !ok { + return + } + s.run(task.prober, task.accountID) + } +} + +func (s *grokImportProbeScheduler) nextTask() (grokImportProbeTask, bool) { + s.mu.Lock() + defer s.mu.Unlock() + if len(s.queue) == 0 { + s.workers-- + return grokImportProbeTask{}, false + } + task := s.queue[0] + s.queue[0] = grokImportProbeTask{} + s.queue = s.queue[1:] + if len(s.queue) == 0 { + s.queue = nil + } + return task, true +} + +func (s *grokImportProbeScheduler) run(prober grokUsageProber, accountID int64) { + defer func() { + if recovered := recover(); recovered != nil { + slog.Error( + "grok_import_active_probe_panic", + "account_id", accountID, + "recovery_type", panicType(recovered), + ) + } + }() + + // Queue time is intentionally excluded: every imported account is probed, + // while this timeout only bounds the actual upstream probe execution. + ctx, cancel := context.WithTimeout(context.Background(), s.timeout) + defer cancel() + result, err := prober.ProbeUsage(ctx, accountID) + if err != nil { + slog.Warn( + "grok_import_active_probe_failed", + "account_id", accountID, + "status", infraerrors.Code(err), + "reason", infraerrors.Reason(err), + ) + return + } + if result == nil { + slog.Warn( + "grok_import_active_probe_failed", + "account_id", accountID, + "reason", "empty_result", + ) + return + } + + slog.Info( + "grok_import_active_probe_completed", + "account_id", accountID, + "model", result.Model, + "status", result.StatusCode, + "headers_observed", result.HeadersObserved, + ) +} + +func panicType(value any) string { + switch value.(type) { + case string: + return "string" + case error: + return "error" + default: + return "unknown" + } +} + +func (h *AccountHandler) scheduleGrokImportProbe(account *service.Account) { + if h == nil { + return + } + defaultGrokImportProbeScheduler.schedule(h.grokImportProber, account) +} + +func (h *GrokOAuthHandler) scheduleGrokImportProbe(account *service.Account) { + if h == nil { + return + } + defaultGrokImportProbeScheduler.schedule(h.importProber, account) +} + +// ProvideAccountHandler injects the Grok active prober for production while +// keeping NewAccountHandler convenient for focused unit tests. +func ProvideAccountHandler( + adminService service.AdminService, + oauthService *service.OAuthService, + openaiOAuthService *service.OpenAIOAuthService, + geminiOAuthService *service.GeminiOAuthService, + antigravityOAuthService *service.AntigravityOAuthService, + rateLimitService *service.RateLimitService, + accountUsageService *service.AccountUsageService, + accountTestService *service.AccountTestService, + concurrencyService *service.ConcurrencyService, + crsSyncService *service.CRSSyncService, + sessionLimitCache service.SessionLimitCache, + rpmCache service.RPMCache, + tokenCacheInvalidator service.TokenCacheInvalidator, + grokQuotaService *service.GrokQuotaService, +) *AccountHandler { + handler := NewAccountHandler( + adminService, + oauthService, + openaiOAuthService, + geminiOAuthService, + antigravityOAuthService, + rateLimitService, + accountUsageService, + accountTestService, + concurrencyService, + crsSyncService, + sessionLimitCache, + rpmCache, + tokenCacheInvalidator, + ) + handler.grokImportProber = grokQuotaService + return handler +} diff --git a/backend/internal/handler/admin/grok_import_probe_handler_test.go b/backend/internal/handler/admin/grok_import_probe_handler_test.go new file mode 100644 index 0000000000..489a13ae6d --- /dev/null +++ b/backend/internal/handler/admin/grok_import_probe_handler_test.go @@ -0,0 +1,119 @@ +//go:build unit + +package admin + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type grokImportAdminService struct { + *stubAdminService + mu sync.Mutex + nextID int64 +} + +func newGrokImportAdminService() *grokImportAdminService { + return &grokImportAdminService{ + stubAdminService: newStubAdminService(), + nextID: 500, + } +} + +func (s *grokImportAdminService) CreateAccount(_ context.Context, input *service.CreateAccountInput) (*service.Account, error) { + s.mu.Lock() + s.nextID++ + id := s.nextID + s.mu.Unlock() + return &service.Account{ + ID: id, + Name: input.Name, + Platform: input.Platform, + Type: input.Type, + Credentials: input.Credentials, + Extra: input.Extra, + ProxyID: input.ProxyID, + Concurrency: input.Concurrency, + Status: service.StatusActive, + Schedulable: true, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + }, nil +} + +type grokImportOAuthClientStub struct{} + +func (grokImportOAuthClientStub) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func (grokImportOAuthClientStub) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func (grokImportOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func TestGrokSSOBatchImportKeepsCreatedAccountsWhenOneAutomaticProbeFails(t *testing.T) { + gin.SetMode(gin.TestMode) + adminService := newGrokImportAdminService() + oauthService := service.NewGrokOAuthService(nil, grokImportOAuthClientStub{}) + defer oauthService.Stop() + prober := newGrokImportProbeStub(3) + prober.failures[502] = infraerrors.New(502, "GROK_TEST_PROBE_FAILED", "sensitive-upstream-body") + handler := NewGrokOAuthHandler(oauthService, adminService, nil) + handler.importProber = prober + + router := gin.New() + router.POST("/api/v1/admin/grok/sso-to-oauth", handler.CreateAccountsFromSSO) + recorder := httptest.NewRecorder() + request := httptest.NewRequest( + http.MethodPost, + "/api/v1/admin/grok/sso-to-oauth", + strings.NewReader(`{"sso_tokens":["sso-one","sso-two","sso-three"]}`), + ) + request.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Contains(t, recorder.Body.String(), `"created"`) + require.NotContains(t, recorder.Body.String(), `GROK_TEST_PROBE_FAILED`) + for i := 0; i < 3; i++ { + awaitGrokProbeSignal(t, prober.done) + } + calls, _, _ := prober.snapshot() + require.Equal(t, map[int64]int{501: 1, 502: 1, 503: 1}, calls) +} + +func TestAccountCreateWithoutAutomaticGrokProbeServiceStillSucceeds(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewAccountHandler( + newGrokImportAdminService(), + nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + ) + + router := gin.New() + router.POST("/api/v1/admin/accounts", handler.Create) + recorder := httptest.NewRecorder() + request := httptest.NewRequest( + http.MethodPost, + "/api/v1/admin/accounts", + strings.NewReader(`{"name":"grok-rt","platform":"grok","type":"oauth","credentials":{"refresh_token":"secret"}}`), + ) + request.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) +} diff --git a/backend/internal/handler/admin/grok_import_probe_test.go b/backend/internal/handler/admin/grok_import_probe_test.go new file mode 100644 index 0000000000..3b8fc0ca6e --- /dev/null +++ b/backend/internal/handler/admin/grok_import_probe_test.go @@ -0,0 +1,232 @@ +//go:build unit + +package admin + +import ( + "bytes" + "context" + "log/slog" + "sync" + "testing" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +type grokImportProbeStub struct { + mu sync.Mutex + calls map[int64]int + failures map[int64]error + active int + maxActive int + deadlineSeen bool + block <-chan struct{} + started chan int64 + done chan int64 +} + +func newGrokImportProbeStub(buffer int) *grokImportProbeStub { + return &grokImportProbeStub{ + calls: make(map[int64]int), + failures: make(map[int64]error), + started: make(chan int64, buffer), + done: make(chan int64, buffer), + } +} + +func (s *grokImportProbeStub) ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) { + _, deadlineSeen := ctx.Deadline() + s.mu.Lock() + s.calls[accountID]++ + s.active++ + if s.active > s.maxActive { + s.maxActive = s.active + } + s.deadlineSeen = s.deadlineSeen || deadlineSeen + s.mu.Unlock() + + s.started <- accountID + var ctxErr error + if s.block != nil { + select { + case <-s.block: + case <-ctx.Done(): + ctxErr = ctx.Err() + } + } + + s.mu.Lock() + s.active-- + failure := s.failures[accountID] + s.mu.Unlock() + s.done <- accountID + if ctxErr != nil { + return nil, ctxErr + } + if failure != nil { + return nil, failure + } + return &service.GrokQuotaProbeResult{ + Source: "active_probe", + Model: "grok-4.5", + StatusCode: 200, + ResetSupported: false, + }, nil +} + +func (s *grokImportProbeStub) snapshot() (map[int64]int, int, bool) { + s.mu.Lock() + defer s.mu.Unlock() + calls := make(map[int64]int, len(s.calls)) + for id, count := range s.calls { + calls[id] = count + } + return calls, s.maxActive, s.deadlineSeen +} + +type grokImportProbeSchedulerTestSnapshot struct { + queued int + workers int + maxWorkers int +} + +func snapshotGrokImportProbeScheduler(s *grokImportProbeScheduler) grokImportProbeSchedulerTestSnapshot { + if s == nil { + return grokImportProbeSchedulerTestSnapshot{} + } + s.mu.Lock() + defer s.mu.Unlock() + return grokImportProbeSchedulerTestSnapshot{ + queued: len(s.queue), + workers: s.workers, + maxWorkers: s.maxWorkers, + } +} + +func newGrokOAuthImportAccount(id int64) *service.Account { + return &service.Account{ + ID: id, + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + } +} + +func awaitGrokProbeSignal(t *testing.T, signals <-chan int64) int64 { + t.Helper() + select { + case id := <-signals: + return id + case <-time.After(time.Second): + t.Fatal("timed out waiting for Grok import probe") + return 0 + } +} + +func TestGrokImportProbeSchedulerProbesSingleAccountOnce(t *testing.T) { + scheduler := newGrokImportProbeScheduler(1, time.Second) + prober := newGrokImportProbeStub(1) + + scheduler.schedule(prober, newGrokOAuthImportAccount(101)) + require.Equal(t, int64(101), awaitGrokProbeSignal(t, prober.done)) + + calls, maxActive, deadlineSeen := prober.snapshot() + require.Equal(t, map[int64]int{101: 1}, calls) + require.Equal(t, 1, maxActive) + require.True(t, deadlineSeen) + require.Eventually(t, func() bool { + snapshot := snapshotGrokImportProbeScheduler(scheduler) + return snapshot.queued == 0 && snapshot.workers == 0 + }, time.Second, 10*time.Millisecond) +} + +func TestGrokImportProbeSchedulerQueuesBatchWithoutPerTaskGoroutines(t *testing.T) { + const taskCount = 100 + release := make(chan struct{}) + scheduler := newGrokImportProbeScheduler(3, time.Second) + prober := newGrokImportProbeStub(taskCount) + prober.block = release + prober.failures[150] = infraerrors.New(502, "GROK_TEST_PROBE_FAILED", "sensitive-upstream-body") + + for id := int64(101); id < 101+taskCount; id++ { + scheduler.schedule(prober, newGrokOAuthImportAccount(id)) + } + for i := 0; i < 3; i++ { + awaitGrokProbeSignal(t, prober.started) + } + snapshot := snapshotGrokImportProbeScheduler(scheduler) + require.Equal(t, 97, snapshot.queued) + require.Equal(t, 3, snapshot.workers) + require.Equal(t, 3, snapshot.maxWorkers) + select { + case id := <-prober.started: + t.Fatalf("probe %d started before a concurrency slot was released", id) + case <-time.After(75 * time.Millisecond): + } + close(release) + for i := 0; i < taskCount; i++ { + awaitGrokProbeSignal(t, prober.done) + } + + calls, maxActive, _ := prober.snapshot() + require.Len(t, calls, taskCount) + for id := int64(101); id < 101+taskCount; id++ { + require.Equal(t, 1, calls[id]) + } + require.Equal(t, 3, maxActive) + require.Eventually(t, func() bool { + snapshot = snapshotGrokImportProbeScheduler(scheduler) + return snapshot.queued == 0 && snapshot.workers == 0 + }, time.Second, 10*time.Millisecond) + require.Equal(t, 3, snapshot.maxWorkers) +} + +func TestGrokImportProbeSchedulerTimeoutCancelsProbe(t *testing.T) { + neverRelease := make(chan struct{}) + scheduler := newGrokImportProbeScheduler(1, 20*time.Millisecond) + prober := newGrokImportProbeStub(1) + prober.block = neverRelease + + scheduler.schedule(prober, newGrokOAuthImportAccount(201)) + require.Equal(t, int64(201), awaitGrokProbeSignal(t, prober.done)) + + calls, _, _ := prober.snapshot() + require.Equal(t, 1, calls[201]) +} + +func TestGrokImportProbeSchedulerSkipsMissingServiceAndNonGrokAccounts(t *testing.T) { + scheduler := newGrokImportProbeScheduler(1, time.Second) + prober := newGrokImportProbeStub(1) + + scheduler.schedule(nil, newGrokOAuthImportAccount(301)) + scheduler.schedule(prober, &service.Account{ID: 302, Platform: service.PlatformOpenAI, Type: service.AccountTypeOAuth}) + scheduler.schedule(prober, &service.Account{ID: 303, Platform: service.PlatformGrok, Type: service.AccountTypeAPIKey}) + + select { + case id := <-prober.started: + t.Fatalf("unexpected probe for account %d", id) + case <-time.After(50 * time.Millisecond): + } + calls, _, _ := prober.snapshot() + require.Empty(t, calls) +} + +func TestGrokImportProbeFailureLogDoesNotIncludeErrorMessage(t *testing.T) { + var logs bytes.Buffer + previousLogger := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + defer slog.SetDefault(previousLogger) + + scheduler := newGrokImportProbeScheduler(1, time.Second) + prober := newGrokImportProbeStub(1) + prober.failures[401] = infraerrors.New(502, "GROK_TEST_PROBE_FAILED", "refresh-token-secret") + scheduler.schedule(prober, newGrokOAuthImportAccount(401)) + awaitGrokProbeSignal(t, prober.done) + + require.Eventually(t, func() bool { + return bytes.Contains(logs.Bytes(), []byte("grok_import_active_probe_failed")) + }, time.Second, 10*time.Millisecond) + require.Contains(t, logs.String(), "GROK_TEST_PROBE_FAILED") + require.NotContains(t, logs.String(), "refresh-token-secret") +} diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go index dafa3076b8..1a309b7c9a 100644 --- a/backend/internal/handler/admin/grok_oauth_handler.go +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -1,20 +1,28 @@ package admin import ( + "context" + "fmt" + "log/slog" "strconv" "strings" + "sync" "github.com/Wei-Shaw/sub2api/internal/handler/dto" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/response" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" ) +const grokSSOImportConcurrency = 3 + type GrokOAuthHandler struct { grokOAuthService *service.GrokOAuthService adminService service.AdminService quotaService *service.GrokQuotaService + importProber grokUsageProber } func NewGrokOAuthHandler( @@ -26,6 +34,7 @@ func NewGrokOAuthHandler( grokOAuthService: grokOAuthService, adminService: adminService, quotaService: quotaService, + importProber: quotaService, } } @@ -202,9 +211,243 @@ func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) { response.ErrorFrom(c, err) return } + h.scheduleGrokImportProbe(account) response.Success(c, dto.AccountFromService(account)) } +type GrokSSOToOAuthRequest struct { + SSOTokens []string `json:"sso_tokens"` + SSOToken string `json:"sso_token"` + Name string `json:"name"` + Notes *string `json:"notes"` + ProxyID *int64 `json:"proxy_id"` + GroupIDs []int64 `json:"group_ids"` + Credentials map[string]any `json:"credentials"` + Extra map[string]any `json:"extra"` + Concurrency int `json:"concurrency"` + LoadFactor *int `json:"load_factor"` + Priority int `json:"priority"` + RateMultiplier *float64 `json:"rate_multiplier"` + ExpiresAt *int64 `json:"expires_at"` + AutoPauseOnExpired *bool `json:"auto_pause_on_expired"` +} + +type GrokSSOToOAuthItemResult struct { + Index int `json:"index"` + Name string `json:"name,omitempty"` + Email string `json:"email,omitempty"` + Account *dto.Account `json:"account,omitempty"` + Error string `json:"error,omitempty"` +} + +type GrokSSOToOAuthResponse struct { + Created []GrokSSOToOAuthItemResult `json:"created"` + Failed []GrokSSOToOAuthItemResult `json:"failed"` +} + +type grokSSOImportJob struct { + index int + token string +} + +type grokSSOImportWorkerResult struct { + created bool + item GrokSSOToOAuthItemResult +} + +func (h *GrokOAuthHandler) CreateAccountsFromSSO(c *gin.Context) { + var req GrokSSOToOAuthRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokens := normalizeSSOImportTokens(req.SSOTokens, req.SSOToken) + if len(tokens) == 0 { + response.BadRequest(c, "sso_tokens is required") + return + } + + ctx := c.Request.Context() + workerCount := grokSSOImportConcurrency + if len(tokens) < workerCount { + workerCount = len(tokens) + } + jobs := make(chan grokSSOImportJob) + items := make([]grokSSOImportWorkerResult, len(tokens)) + var wg sync.WaitGroup + for i := 0; i < workerCount; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for job := range jobs { + items[job.index] = h.safeCreateAccountFromSSOToken(ctx, req, job.token, job.index+1, len(tokens)) + } + }() + } + for i, token := range tokens { + jobs <- grokSSOImportJob{index: i, token: token} + } + close(jobs) + wg.Wait() + + result := GrokSSOToOAuthResponse{ + Created: make([]GrokSSOToOAuthItemResult, 0, len(tokens)), + Failed: make([]GrokSSOToOAuthItemResult, 0), + } + for _, item := range items { + if item.created { + result.Created = append(result.Created, item.item) + } else { + result.Failed = append(result.Failed, item.item) + } + } + response.Success(c, result) +} + +func (h *GrokOAuthHandler) safeCreateAccountFromSSOToken(ctx context.Context, req GrokSSOToOAuthRequest, token string, index, total int) (result grokSSOImportWorkerResult) { + defer func() { + if recovered := recover(); recovered != nil { + slog.Error("grok_sso_import_worker_panic", "index", index, "recover", recovered) + result = grokSSOImportWorkerResult{ + item: GrokSSOToOAuthItemResult{ + Index: index, + Error: fmt.Sprintf("internal worker panic: %v", recovered), + }, + } + } + }() + return h.createAccountFromSSOToken(ctx, req, token, index, total) +} + +func (h *GrokOAuthHandler) createAccountFromSSOToken(ctx context.Context, req GrokSSOToOAuthRequest, token string, index, total int) grokSSOImportWorkerResult { + tokenInfo, err := h.grokOAuthService.ConvertFromSSO(ctx, token, req.ProxyID) + if err != nil { + return grokSSOImportWorkerResult{item: GrokSSOToOAuthItemResult{Index: index, Error: grokSSOImportErrorMessage(err)}} + } + + credentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo) + credentials = service.MergeCredentials(cloneGrokSSOMap(req.Credentials), credentials) + name := grokSSOImportAccountName(req.Name, tokenInfo, index, total) + expiresAt, autoPauseOnExpired := grokSSOImportExpiry(req.ExpiresAt, req.AutoPauseOnExpired, tokenInfo) + account, err := h.adminService.CreateAccount(ctx, &service.CreateAccountInput{ + Name: name, + Notes: req.Notes, + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + Credentials: credentials, + Extra: cloneGrokSSOMap(req.Extra), + ProxyID: req.ProxyID, + Concurrency: req.Concurrency, + LoadFactor: req.LoadFactor, + Priority: req.Priority, + RateMultiplier: req.RateMultiplier, + GroupIDs: append([]int64(nil), req.GroupIDs...), + ExpiresAt: expiresAt, + AutoPauseOnExpired: autoPauseOnExpired, + }) + if err != nil { + return grokSSOImportWorkerResult{item: GrokSSOToOAuthItemResult{Index: index, Name: name, Email: tokenInfo.Email, Error: grokSSOImportErrorMessage(err)}} + } + h.scheduleGrokImportProbe(account) + return grokSSOImportWorkerResult{ + created: true, + item: GrokSSOToOAuthItemResult{ + Index: index, + Name: name, + Email: tokenInfo.Email, + Account: dto.AccountFromService(account), + }, + } +} + +func grokSSOImportExpiry(requestExpiresAt *int64, requestAutoPause *bool, tokenInfo *service.GrokTokenInfo) (*int64, *bool) { + if tokenInfo == nil || strings.TrimSpace(tokenInfo.RefreshToken) != "" || tokenInfo.ExpiresAt <= 0 { + return requestExpiresAt, requestAutoPause + } + + expiresAt := tokenInfo.ExpiresAt + if requestExpiresAt != nil && *requestExpiresAt > 0 && *requestExpiresAt < expiresAt { + expiresAt = *requestExpiresAt + } + autoPause := true + return &expiresAt, &autoPause +} + +func cloneGrokSSOMap(source map[string]any) map[string]any { + if source == nil { + return nil + } + clone := make(map[string]any, len(source)) + for key, value := range source { + clone[key] = cloneGrokSSOValue(value) + } + return clone +} + +func cloneGrokSSOValue(value any) any { + switch v := value.(type) { + case map[string]any: + return cloneGrokSSOMap(v) + case []any: + clone := make([]any, len(v)) + for i, item := range v { + clone[i] = cloneGrokSSOValue(item) + } + return clone + default: + return value + } +} + +func normalizeSSOImportTokens(tokens []string, single string) []string { + items := make([]string, 0, len(tokens)+1) + if strings.TrimSpace(single) != "" { + items = append(items, single) + } + items = append(items, tokens...) + seen := make(map[string]struct{}, len(items)) + result := make([]string, 0, len(items)) + for _, item := range items { + parts := strings.Split(strings.NewReplacer(",", "\n", "\r", "\n").Replace(item), "\n") + for _, token := range parts { + if token = xai.NormalizeSSOToken(token); token == "" { + continue + } + if _, ok := seen[token]; ok { + continue + } + seen[token] = struct{}{} + result = append(result, token) + } + } + return result +} + +func grokSSOImportAccountName(base string, tokenInfo *service.GrokTokenInfo, index, total int) string { + base = strings.TrimSpace(base) + if base == "" && tokenInfo != nil { + base = strings.TrimSpace(tokenInfo.Email) + } + if base == "" { + base = "Grok OAuth Account" + } + if total > 1 { + return base + " #" + strconv.Itoa(index) + } + return base +} + +func grokSSOImportErrorMessage(err error) string { + status := infraerrors.FromError(err) + if status == nil { + return "" + } + if status.Reason != "" { + return status.Reason + ": " + status.Message + } + return status.Message +} + func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) { accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil { @@ -215,7 +458,7 @@ func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) { response.BadRequest(c, "grok quota service is not enabled") return } - result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID) + result, err := h.quotaService.QueryQuota(c.Request.Context(), accountID) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 6ac77e0e56..64ea044aa3 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -8,6 +8,7 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" "time" @@ -41,17 +42,35 @@ func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, u } type grokQuotaHandlerUpstream struct { - resp *http.Response - lastReq *http.Request - lastBody []byte + mu sync.Mutex + requests []*http.Request + bodies [][]byte } func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { - u.lastReq = req + var body []byte if req.Body != nil { - u.lastBody, _ = io.ReadAll(req.Body) + body, _ = io.ReadAll(req.Body) } - return u.resp, nil + u.mu.Lock() + u.requests = append(u.requests, req) + u.bodies = append(u.bodies, body) + u.mu.Unlock() + if req.URL.Path == "/v1/responses" { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"8"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }, nil + } + payload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}` + if req.URL.RawQuery == "format=credits" { + payload = `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}}}` + } + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil } func (u *grokQuotaHandlerUpstream) DoWithTLS( @@ -77,14 +96,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), }, }} - upstream := &grokQuotaHandlerUpstream{resp: &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{ - "X-Ratelimit-Limit-Requests": []string{"10"}, - "X-Ratelimit-Remaining-Requests": []string{"8"}, - }, - Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), - }} + upstream := &grokQuotaHandlerUpstream{} quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream) handler := NewGrokOAuthHandler(nil, nil, quotaService) @@ -95,12 +107,23 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { router.ServeHTTP(rec, req) require.Equal(t, http.StatusOK, rec.Code) - require.Contains(t, rec.Body.String(), `"source":"active_probe"`) + require.Contains(t, rec.Body.String(), `"source":"hybrid_probe"`) + require.Contains(t, rec.Body.String(), `"billing":`) + require.Contains(t, rec.Body.String(), `"snapshot":`) require.Contains(t, rec.Body.String(), `"headers_observed":true`) require.NotContains(t, rec.Body.String(), "access-token") - require.Equal(t, xai.DefaultBaseURL+"/responses", upstream.lastReq.URL.String()) - require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Contains(t, string(upstream.lastBody), `"store":false`) + upstream.mu.Lock() + requests := append([]*http.Request(nil), upstream.requests...) + bodies := append([][]byte(nil), upstream.bodies...) + upstream.mu.Unlock() + require.Len(t, requests, 3) + for i, upstreamReq := range requests { + require.Equal(t, "Bearer access-token", upstreamReq.Header.Get("Authorization")) + if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/responses" { + require.Contains(t, string(bodies[i]), `"model":"grok-4.5"`) + require.Contains(t, string(bodies[i]), `"store":false`) + } + } require.NotNil(t, repo.updates[42]) } @@ -145,3 +168,51 @@ func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) { require.NotContains(t, rec.Body.String(), "secret") require.NotContains(t, rec.Body.String(), "client-secret-like-value") } + +func TestGrokSSOImportExpiryUsesTokenExpiryWithoutRefreshToken(t *testing.T) { + tokenExpiry := time.Now().Add(6 * time.Hour).Unix() + expiresAt, autoPause := grokSSOImportExpiry(nil, nil, &service.GrokTokenInfo{ + ExpiresAt: tokenExpiry, + }) + + require.NotNil(t, expiresAt) + require.Equal(t, tokenExpiry, *expiresAt) + require.NotNil(t, autoPause) + require.True(t, *autoPause) +} + +func TestGrokSSOImportExpiryUsesEarlierRequestedExpiryWithoutRefreshToken(t *testing.T) { + requestedExpiry := time.Now().Add(2 * time.Hour).Unix() + tokenExpiry := time.Now().Add(6 * time.Hour).Unix() + requestedAutoPause := false + expiresAt, autoPause := grokSSOImportExpiry(&requestedExpiry, &requestedAutoPause, &service.GrokTokenInfo{ + ExpiresAt: tokenExpiry, + }) + + require.NotNil(t, expiresAt) + require.Equal(t, requestedExpiry, *expiresAt) + require.NotNil(t, autoPause) + require.True(t, *autoPause) +} + +func TestGrokSSOImportExpiryPreservesRequestSettingsWithRefreshToken(t *testing.T) { + requestedExpiry := time.Now().Add(2 * time.Hour).Unix() + requestedAutoPause := false + expiresAt, autoPause := grokSSOImportExpiry(&requestedExpiry, &requestedAutoPause, &service.GrokTokenInfo{ + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(6 * time.Hour).Unix(), + }) + + require.Same(t, &requestedExpiry, expiresAt) + require.Same(t, &requestedAutoPause, autoPause) +} + +func TestGrokSSOImportWorkerRecoversPanic(t *testing.T) { + h := &GrokOAuthHandler{} + result := h.safeCreateAccountFromSSOToken(context.Background(), GrokSSOToOAuthRequest{}, "token", 2, 3) + // Without a service, createAccountFromSSOToken would panic on nil service access. + // Recovery must convert that into a failed item and keep the worker alive. + require.False(t, result.created) + require.Equal(t, 2, result.item.Index) + require.Contains(t, result.item.Error, "internal worker panic") +} diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 56a0b29ed0..02a4446846 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -110,6 +110,7 @@ type CreateGroupRequest struct { VideoPrice480P *float64 `json:"video_price_480p"` VideoPrice720P *float64 `json:"video_price_720p"` VideoPrice1080P *float64 `json:"video_price_1080p"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` ClaudeCodeOnly bool `json:"claude_code_only"` FallbackGroupID *int64 `json:"fallback_group_id"` FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` @@ -163,6 +164,7 @@ type UpdateGroupRequest struct { VideoPrice480P *float64 `json:"video_price_480p"` VideoPrice720P *float64 `json:"video_price_720p"` VideoPrice1080P *float64 `json:"video_price_1080p"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` ClaudeCodeOnly *bool `json:"claude_code_only"` FallbackGroupID *int64 `json:"fallback_group_id"` FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` @@ -334,6 +336,7 @@ func (h *GroupHandler) Create(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + WebSearchPricePerCall: req.WebSearchPricePerCall, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest, @@ -402,6 +405,7 @@ func (h *GroupHandler) Update(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + WebSearchPricePerCall: req.WebSearchPricePerCall, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest, diff --git a/backend/internal/handler/admin/openai_oauth_handler.go b/backend/internal/handler/admin/openai_oauth_handler.go index d7a756bd00..78d57299b6 100644 --- a/backend/internal/handler/admin/openai_oauth_handler.go +++ b/backend/internal/handler/admin/openai_oauth_handler.go @@ -304,6 +304,10 @@ func (h *OpenAIOAuthHandler) CreateAccountFromCodexPAT(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.Concurrency != nil && *req.Concurrency < 0 { response.BadRequest(c, "concurrency must be >= 0") return diff --git a/backend/internal/handler/admin/ops_system_log_handler.go b/backend/internal/handler/admin/ops_system_log_handler.go index 9f3c8b893a..1b6af45976 100644 --- a/backend/internal/handler/admin/ops_system_log_handler.go +++ b/backend/internal/handler/admin/ops_system_log_handler.go @@ -15,6 +15,7 @@ import ( type opsSystemLogCleanupRequest struct { StartTime string `json:"start_time"` EndTime string `json:"end_time"` + Host string `json:"host"` Level string `json:"level"` Component string `json:"component"` @@ -56,6 +57,7 @@ func (h *OpsHandler) ListSystemLogs(c *gin.Context) { PageSize: pageSize, StartTime: &start, EndTime: &end, + Host: strings.TrimSpace(c.Query("host")), Level: strings.TrimSpace(c.Query("level")), Component: strings.TrimSpace(c.Query("component")), RequestID: strings.TrimSpace(c.Query("request_id")), @@ -153,6 +155,7 @@ func (h *OpsHandler) CleanupSystemLogs(c *gin.Context) { filter := &service.OpsSystemLogCleanupFilter{ StartTime: start, EndTime: end, + Host: strings.TrimSpace(req.Host), Level: strings.TrimSpace(req.Level), Component: strings.TrimSpace(req.Component), RequestID: strings.TrimSpace(req.RequestID), diff --git a/backend/internal/handler/admin/ops_system_log_handler_test.go b/backend/internal/handler/admin/ops_system_log_handler_test.go index 9557fce442..3390fbe3cb 100644 --- a/backend/internal/handler/admin/ops_system_log_handler_test.go +++ b/backend/internal/handler/admin/ops_system_log_handler_test.go @@ -2,6 +2,7 @@ package admin import ( "bytes" + "context" "encoding/json" "net/http" "net/http/httptest" @@ -19,6 +20,26 @@ type responseEnvelope struct { Data json.RawMessage `json:"data"` } +type opsSystemLogCaptureRepo struct { + service.OpsRepository + listFilter *service.OpsSystemLogFilter + cleanupFilter *service.OpsSystemLogCleanupFilter +} + +func (r *opsSystemLogCaptureRepo) ListSystemLogs(_ context.Context, filter *service.OpsSystemLogFilter) (*service.OpsSystemLogList, error) { + r.listFilter = filter + return &service.OpsSystemLogList{Logs: []*service.OpsSystemLog{}, Page: filter.Page, PageSize: filter.PageSize}, nil +} + +func (r *opsSystemLogCaptureRepo) DeleteSystemLogs(_ context.Context, filter *service.OpsSystemLogCleanupFilter) (int64, error) { + r.cleanupFilter = filter + return 1, nil +} + +func (r *opsSystemLogCaptureRepo) InsertSystemLogCleanupAudit(_ context.Context, _ *service.OpsSystemLogCleanupAudit) error { + return nil +} + func newOpsSystemLogTestRouter(handler *OpsHandler, withUser bool) *gin.Engine { gin.SetMode(gin.TestMode) r := gin.New() @@ -121,6 +142,23 @@ func TestOpsSystemLogHandler_ListSuccess(t *testing.T) { } } +func TestOpsSystemLogHandler_ListAcceptsHost(t *testing.T) { + repo := &opsSystemLogCaptureRepo{} + svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + h := NewOpsHandler(svc) + r := newOpsSystemLogTestRouter(h, false) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/logs?host=api-node-1", nil) + r.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("status=%d, want 200", w.Code) + } + if repo.listFilter == nil || repo.listFilter.Host != "api-node-1" { + t.Fatalf("host filter = %+v, want api-node-1", repo.listFilter) + } +} + func TestOpsSystemLogHandler_CleanupUnauthorized(t *testing.T) { svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) h := NewOpsHandler(svc) @@ -205,6 +243,24 @@ func TestOpsSystemLogHandler_CleanupAcceptsAPIKeyID(t *testing.T) { } } +func TestOpsSystemLogHandler_CleanupAcceptsHost(t *testing.T) { + repo := &opsSystemLogCaptureRepo{} + svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + h := NewOpsHandler(svc) + r := newOpsSystemLogTestRouter(h, true) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/logs/cleanup", bytes.NewBufferString(`{"host":"api-node-1"}`)) + req.Header.Set("Content-Type", "application/json") + r.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("status=%d, want 200", w.Code) + } + if repo.cleanupFilter == nil || repo.cleanupFilter.Host != "api-node-1" { + t.Fatalf("host filter = %+v, want api-node-1", repo.cleanupFilter) + } +} + func TestOpsSystemLogHandler_CleanupInvalidAPIKeyID(t *testing.T) { svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) h := NewOpsHandler(svc) diff --git a/backend/internal/handler/admin/ops_ws_handler.go b/backend/internal/handler/admin/ops_ws_handler.go index 75fd7ea002..e4c42cc9c0 100644 --- a/backend/internal/handler/admin/ops_ws_handler.go +++ b/backend/internal/handler/admin/ops_ws_handler.go @@ -16,6 +16,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" @@ -323,7 +324,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) { // If realtime monitoring is disabled, prefer a successful WS upgrade followed by a clean close // with a deterministic close code. This prevents clients from spinning on 404/1006 reconnect loops. if !h.opsService.IsRealtimeMonitoringEnabled(c.Request.Context()) { - conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c)) if err != nil { c.JSON(http.StatusNotFound, gin.H{"error": "ops realtime monitoring is disabled"}) return @@ -358,7 +359,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) { defer releaseOpsWSIPSlot(clientIP) } - conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c)) if err != nil { logger.LegacyPrintf("handler.admin.ops_ws", "[OpsWS] upgrade failed: %v", err) return diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 6270c2b982..3c45c3b95e 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -199,6 +199,7 @@ func groupFromServiceBase(g *service.Group) Group { VideoPrice480P: g.VideoPrice480P, VideoPrice720P: g.VideoPrice720P, VideoPrice1080P: g.VideoPrice1080P, + WebSearchPricePerCall: g.WebSearchPricePerCall, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, FallbackGroupIDOnInvalidRequest: g.FallbackGroupIDOnInvalidRequest, @@ -598,54 +599,55 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog { requestedModel = l.Model } return UsageLog{ - ID: l.ID, - UserID: l.UserID, - APIKeyID: l.APIKeyID, - AccountID: l.AccountID, - RequestID: l.RequestID, - Model: requestedModel, - ServiceTier: l.ServiceTier, - ReasoningEffort: l.ReasoningEffort, - InboundEndpoint: l.InboundEndpoint, - GroupID: l.GroupID, - SubscriptionID: l.SubscriptionID, - InputTokens: l.InputTokens, - OutputTokens: l.OutputTokens, - CacheCreationTokens: l.CacheCreationTokens, - CacheReadTokens: l.CacheReadTokens, - CacheCreation5mTokens: l.CacheCreation5mTokens, - CacheCreation1hTokens: l.CacheCreation1hTokens, - InputCost: l.InputCost, - OutputCost: l.OutputCost, - CacheCreationCost: l.CacheCreationCost, - CacheReadCost: l.CacheReadCost, - TotalCost: l.TotalCost, - ActualCost: l.ActualCost, - RateMultiplier: l.RateMultiplier, - BillingType: l.BillingType, - RequestType: requestType.String(), - Stream: stream, - OpenAIWSMode: openAIWSMode, - DurationMs: l.DurationMs, - FirstTokenMs: l.FirstTokenMs, - ImageCount: l.ImageCount, - ImageSize: l.ImageSize, - ImageInputSize: l.ImageInputSize, - ImageOutputSize: l.ImageOutputSize, - ImageOutputTokens: l.ImageOutputTokens, - ImageOutputCost: l.ImageOutputCost, - ImageSizeSource: l.ImageSizeSource, - ImageSizeBreakdown: l.ImageSizeBreakdown, - MediaType: l.MediaType, - UserAgent: l.UserAgent, - IPAddress: l.IPAddress, - CacheTTLOverridden: l.CacheTTLOverridden, - BillingMode: l.BillingMode, - CreatedAt: l.CreatedAt, - User: UserFromServiceShallow(l.User), - APIKey: APIKeyFromService(l.APIKey), - Group: GroupFromServiceShallow(l.Group), - Subscription: UserSubscriptionFromService(l.Subscription), + ID: l.ID, + UserID: l.UserID, + APIKeyID: l.APIKeyID, + AccountID: l.AccountID, + RequestID: l.RequestID, + Model: requestedModel, + ServiceTier: l.ServiceTier, + ReasoningEffort: l.ReasoningEffort, + InboundEndpoint: l.InboundEndpoint, + GroupID: l.GroupID, + SubscriptionID: l.SubscriptionID, + InputTokens: l.InputTokens, + OutputTokens: l.OutputTokens, + CacheCreationTokens: l.CacheCreationTokens, + CacheReadTokens: l.CacheReadTokens, + CacheCreation5mTokens: l.CacheCreation5mTokens, + CacheCreation1hTokens: l.CacheCreation1hTokens, + InputCost: l.InputCost, + OutputCost: l.OutputCost, + CacheCreationCost: l.CacheCreationCost, + CacheReadCost: l.CacheReadCost, + TotalCost: l.TotalCost, + ActualCost: l.ActualCost, + RateMultiplier: l.RateMultiplier, + LongContextBillingApplied: l.LongContextBillingApplied, + BillingType: l.BillingType, + RequestType: requestType.String(), + Stream: stream, + OpenAIWSMode: openAIWSMode, + DurationMs: l.DurationMs, + FirstTokenMs: l.FirstTokenMs, + ImageCount: l.ImageCount, + ImageSize: l.ImageSize, + ImageInputSize: l.ImageInputSize, + ImageOutputSize: l.ImageOutputSize, + ImageOutputTokens: l.ImageOutputTokens, + ImageOutputCost: l.ImageOutputCost, + ImageSizeSource: l.ImageSizeSource, + ImageSizeBreakdown: l.ImageSizeBreakdown, + MediaType: l.MediaType, + UserAgent: l.UserAgent, + IPAddress: l.IPAddress, + CacheTTLOverridden: l.CacheTTLOverridden, + BillingMode: l.BillingMode, + CreatedAt: l.CreatedAt, + User: UserFromServiceShallow(l.User), + APIKey: APIKeyFromService(l.APIKey), + Group: GroupFromServiceShallow(l.Group), + Subscription: UserSubscriptionFromService(l.Subscription), } } diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 7cfd102880..619926c1e4 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -120,6 +120,8 @@ type Group struct { VideoPrice480P *float64 `json:"video_price_480p"` VideoPrice720P *float64 `json:"video_price_720p"` VideoPrice1080P *float64 `json:"video_price_1080p"` + // Codex alpha/search 网页搜索单次价格(USD/次);null 表示使用默认价 0.01 + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` // Claude Code 客户端限制 ClaudeCodeOnly bool `json:"claude_code_only"` @@ -480,13 +482,14 @@ type UsageLog struct { CacheCreation5mTokens int `json:"cache_creation_5m_tokens"` CacheCreation1hTokens int `json:"cache_creation_1h_tokens"` - InputCost float64 `json:"input_cost"` - OutputCost float64 `json:"output_cost"` - CacheCreationCost float64 `json:"cache_creation_cost"` - CacheReadCost float64 `json:"cache_read_cost"` - TotalCost float64 `json:"total_cost"` - ActualCost float64 `json:"actual_cost"` - RateMultiplier float64 `json:"rate_multiplier"` + InputCost float64 `json:"input_cost"` + OutputCost float64 `json:"output_cost"` + CacheCreationCost float64 `json:"cache_creation_cost"` + CacheReadCost float64 `json:"cache_read_cost"` + TotalCost float64 `json:"total_cost"` + ActualCost float64 `json:"actual_cost"` + RateMultiplier float64 `json:"rate_multiplier"` + LongContextBillingApplied bool `json:"long_context_billing_applied"` BillingType int8 `json:"billing_type"` RequestType string `json:"request_type"` diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 0b9930c5cc..5e4d84ba72 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -18,11 +18,14 @@ const ( EndpointMessages = "/v1/messages" EndpointChatCompletions = "/v1/chat/completions" EndpointEmbeddings = "/v1/embeddings" + EndpointAlphaSearch = "/v1/alpha/search" EndpointResponses = "/v1/responses" EndpointResponsesCompact = "/v1/responses/compact" EndpointImagesGenerations = "/v1/images/generations" EndpointImagesEdits = "/v1/images/edits" EndpointVideosGenerations = "/v1/videos/generations" + EndpointVideosEdits = "/v1/videos/edits" + EndpointVideosExtensions = "/v1/videos/extensions" EndpointVideos = "/v1/videos" EndpointGeminiModels = "/v1beta/models" ) @@ -75,6 +78,8 @@ func NormalizeInboundEndpoint(path string) string { switch { case strings.Contains(path, EndpointEmbeddings): return EndpointEmbeddings + case strings.Contains(path, EndpointAlphaSearch) || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/alpha/search") || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/backend-api/codex/alpha/search"): + return EndpointAlphaSearch case strings.Contains(path, EndpointChatCompletions): return EndpointChatCompletions case strings.Contains(path, EndpointMessages): @@ -85,6 +90,10 @@ func NormalizeInboundEndpoint(path string) string { return EndpointImagesEdits case strings.Contains(path, EndpointVideosGenerations) || strings.Contains(path, "/videos/generations"): return EndpointVideosGenerations + case strings.Contains(path, EndpointVideosEdits) || strings.Contains(path, "/videos/edits"): + return EndpointVideosEdits + case strings.Contains(path, EndpointVideosExtensions) || strings.Contains(path, "/videos/extensions"): + return EndpointVideosExtensions case strings.Contains(path, EndpointVideos) || strings.Contains(path, "/videos/"): return EndpointVideos case strings.Contains(path, EndpointResponsesCompact) || isResponsesCompactAliasPath(path): @@ -155,8 +164,11 @@ func isBareOrSubpathOf(path, root string) bool { // account platform and the normalized inbound endpoint. // // Platform-specific rules: -// - OpenAI always forwards to /v1/responses (with optional subpath -// such as /v1/responses/compact preserved from the raw URL). +// - OpenAI and Grok text compatibility routes forward to /v1/responses +// (with optional subpath such as /v1/responses/compact preserved from +// the raw URL); native endpoints such as embeddings and alpha search +// retain their paths. Grok raw Chat requests override this through the +// forwarding result consumed by resolveOpenAIUpstreamEndpoint. // - Anthropic → /v1/messages // - Gemini → /v1beta/models // - Antigravity → /v1/messages (Claude) or gemini (Gemini) @@ -167,7 +179,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { switch platform { case service.PlatformOpenAI, service.PlatformGrok: - if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos { + if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideosEdits || inbound == EndpointVideosExtensions || inbound == EndpointVideos { return inbound } // OpenAI forwards everything to the Responses API. diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index 96ed1292b3..dc8ad728c8 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -25,6 +25,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) { {"/v1/messages", EndpointMessages}, {"/v1/chat/completions", EndpointChatCompletions}, {"/v1/embeddings", EndpointEmbeddings}, + {"/v1/alpha/search", EndpointAlphaSearch}, {"/v1/responses", EndpointResponses}, {"/v1/responses/compact", EndpointResponsesCompact}, {"/v1/responses/compact/detail", EndpointResponsesCompact}, @@ -50,11 +51,13 @@ func TestNormalizeInboundEndpoint(t *testing.T) { {"/responses", EndpointResponses}, {"/responses/compact", EndpointResponsesCompact}, {"/responses/compact/detail", EndpointResponsesCompact}, + {"/alpha/search", EndpointAlphaSearch}, // Bare Codex direct alias route — root vs. compact. {"/backend-api/codex/responses", EndpointResponses}, {"/backend-api/codex/responses/compact", EndpointResponsesCompact}, {"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact}, + {"/backend-api/codex/alpha/search", EndpointAlphaSearch}, // Must NOT generalize to arbitrary paths merely ending in // "/responses" (or "/responses/compact") that are unrelated to @@ -119,8 +122,11 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { {"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses}, {"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses}, {"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings}, + {"openai alpha search", EndpointAlphaSearch, "/backend-api/codex/alpha/search", service.PlatformOpenAI, EndpointAlphaSearch}, {"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations}, {"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits}, + {"grok chat defaults to responses without runtime result", EndpointChatCompletions, "/v1/chat/completions", service.PlatformGrok, EndpointResponses}, + {"grok responses", EndpointResponses, "/v1/responses", service.PlatformGrok, EndpointResponses}, {"grok video generations", EndpointVideosGenerations, "/v1/videos/generations", service.PlatformGrok, EndpointVideosGenerations}, {"grok video status", EndpointVideos, "/videos/req_123", service.PlatformGrok, EndpointVideos}, @@ -138,6 +144,59 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { } } +func TestResolveOpenAIUpstreamEndpointPrefersForwardResult(t *testing.T) { + tests := []struct { + name string + account *service.Account + result *service.OpenAIForwardResult + runtimeEndpoint string + want string + }{ + { + name: "grok raw chat result overrides stale context", + account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}, + result: &service.OpenAIForwardResult{UpstreamEndpoint: EndpointChatCompletions}, + runtimeEndpoint: EndpointResponses, + want: EndpointChatCompletions, + }, + { + name: "grok chat bridged to responses", + account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}, + result: &service.OpenAIForwardResult{UpstreamEndpoint: EndpointResponses}, + want: EndpointResponses, + }, + { + name: "grok empty result keeps responses default", + account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}, + result: &service.OpenAIForwardResult{}, + want: EndpointResponses, + }, + { + name: "grok raw error uses runtime endpoint", + account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}, + runtimeEndpoint: EndpointChatCompletions, + want: EndpointChatCompletions, + }, + { + name: "openai behavior remains responses", + account: &service.Account{Platform: service.PlatformOpenAI, Type: service.AccountTypeOAuth}, + result: &service.OpenAIForwardResult{}, + want: EndpointResponses, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, EndpointChatCompletions, nil) + c.Set(ctxKeyInboundEndpoint, EndpointChatCompletions) + service.SetActualOpenAIUpstreamEndpoint(c, tt.runtimeEndpoint) + require.Equal(t, tt.want, resolveOpenAIUpstreamEndpoint(c, tt.account, tt.result)) + }) + } +} + // ────────────────────────────────────────────────────────── // responsesSubpathSuffix // ────────────────────────────────────────────────────────── diff --git a/backend/internal/handler/failover_loop.go b/backend/internal/handler/failover_loop.go index 6d8ddc7236..5838e58f48 100644 --- a/backend/internal/handler/failover_loop.go +++ b/backend/internal/handler/failover_loop.go @@ -29,7 +29,8 @@ const ( ) const ( - // maxSameAccountRetries 同账号重试次数上限(针对 RetryableOnSameAccount 错误) + // maxSameAccountRetries 同账号重试次数默认上限(针对 RetryableOnSameAccount 错误)。 + // 生产调用方通常传入账号级配置 account.GetPoolModeRetryCount(),该常量仅作兜底/测试默认值。 maxSameAccountRetries = 3 // sameAccountRetryDelay 同账号重试间隔 sameAccountRetryDelay = 500 * time.Millisecond @@ -67,6 +68,7 @@ func (s *FailoverState) HandleFailoverError( gatewayService TempUnscheduler, accountID int64, platform string, + retryLimit int, failoverErr *service.UpstreamFailoverError, ) FailoverAction { s.LastFailoverErr = failoverErr @@ -76,14 +78,15 @@ func (s *FailoverState) HandleFailoverError( s.ForceCacheBilling = true } - // 同账号重试:对 RetryableOnSameAccount 的临时性错误,先在同一账号上重试 - if failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < maxSameAccountRetries { + // 同账号重试:对 RetryableOnSameAccount 的临时性错误,先在同一账号上重试。 + // 重试次数上限 retryLimit 由调用方传入(账号级 pool_mode_retry_count 配置)。 + if failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < retryLimit { s.SameAccountRetryCount[accountID]++ logger.FromContext(ctx).Warn("gateway.failover_same_account_retry", zap.Int64("account_id", accountID), zap.Int("upstream_status", failoverErr.StatusCode), zap.Int("same_account_retry_count", s.SameAccountRetryCount[accountID]), - zap.Int("same_account_retry_max", maxSameAccountRetries), + zap.Int("same_account_retry_max", retryLimit), ) if !sleepWithContext(ctx, sameAccountRetryDelay) { return FailoverCanceled diff --git a/backend/internal/handler/failover_loop_test.go b/backend/internal/handler/failover_loop_test.go index 2c65ebc2c8..9fabe75f25 100644 --- a/backend/internal/handler/failover_loop_test.go +++ b/backend/internal/handler/failover_loop_test.go @@ -133,7 +133,7 @@ func TestHandleFailoverError_BasicSwitch(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, false, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SwitchCount) @@ -150,7 +150,7 @@ func TestHandleFailoverError_BasicSwitch(t *testing.T) { err := newTestFailoverErr(500, false, false) start := time.Now() - action := fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformAntigravity, err) + action := fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverContinue, action) @@ -166,7 +166,7 @@ func TestHandleFailoverError_BasicSwitch(t *testing.T) { err := newTestFailoverErr(500, false, false) start := time.Now() - action := fs.HandleFailoverError(context.Background(), mock, 200, service.PlatformAntigravity, err) + action := fs.HandleFailoverError(context.Background(), mock, 200, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverContinue, action) @@ -181,19 +181,19 @@ func TestHandleFailoverError_BasicSwitch(t *testing.T) { // 第一次切换:0→1 err1 := newTestFailoverErr(500, false, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err1) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err1) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SwitchCount) // 第二次切换:1→2 err2 := newTestFailoverErr(502, false, false) - action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", err2) + action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, err2) require.Equal(t, FailoverContinue, action) require.Equal(t, 2, fs.SwitchCount) // 第三次已耗尽:SwitchCount(2) >= MaxSwitches(2) err3 := newTestFailoverErr(503, false, false) - action = fs.HandleFailoverError(context.Background(), mock, 300, "openai", err3) + action = fs.HandleFailoverError(context.Background(), mock, 300, "openai", maxSameAccountRetries, err3) require.Equal(t, FailoverExhausted, action) require.Equal(t, 2, fs.SwitchCount, "耗尽时不应继续递增") @@ -212,7 +212,7 @@ func TestHandleFailoverError_BasicSwitch(t *testing.T) { fs := NewFailoverState(0, false) err := newTestFailoverErr(500, false, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverExhausted, action) require.Equal(t, 0, fs.SwitchCount) require.Contains(t, fs.FailedAccountIDs, int64(100)) @@ -229,7 +229,7 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) { fs := NewFailoverState(3, true) // hasBoundSession=true err := newTestFailoverErr(500, false, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.True(t, fs.ForceCacheBilling) }) @@ -238,7 +238,7 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, false, true) // ForceCacheBilling=true - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.True(t, fs.ForceCacheBilling) }) @@ -247,7 +247,7 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, false, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.False(t, fs.ForceCacheBilling) }) @@ -257,12 +257,12 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) { // 第一次:ForceCacheBilling=true → 设置 err1 := newTestFailoverErr(500, false, true) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err1) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err1) require.True(t, fs.ForceCacheBilling) // 第二次:ForceCacheBilling=false → 仍然保持 true err2 := newTestFailoverErr(502, false, false) - fs.HandleFailoverError(context.Background(), mock, 200, "openai", err2) + fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, err2) require.True(t, fs.ForceCacheBilling, "ForceCacheBilling 一旦设置不应被重置") }) } @@ -278,7 +278,7 @@ func TestHandleFailoverError_SameAccountRetry(t *testing.T) { err := newTestFailoverErr(400, true, false) start := time.Now() - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverContinue, action) @@ -297,7 +297,7 @@ func TestHandleFailoverError_SameAccountRetry(t *testing.T) { err := newTestFailoverErr(400, true, false) for i := 1; i <= maxSameAccountRetries; i++ { - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, i, fs.SameAccountRetryCount[100]) } @@ -311,12 +311,12 @@ func TestHandleFailoverError_SameAccountRetry(t *testing.T) { err := newTestFailoverErr(400, true, false) for i := 0; i < maxSameAccountRetries; i++ { - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) } require.Equal(t, maxSameAccountRetries, fs.SameAccountRetryCount[100]) // 第 maxSameAccountRetries+1 次:重试耗尽,应切换账号 - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SwitchCount) require.Contains(t, fs.FailedAccountIDs, int64(100)) @@ -333,12 +333,12 @@ func TestHandleFailoverError_SameAccountRetry(t *testing.T) { err := newTestFailoverErr(400, true, false) // 账号 100 第一次重试 - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SameAccountRetryCount[100]) // 账号 200 第一次重试(独立计数) - action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", err) + action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SameAccountRetryCount[200]) require.Equal(t, 1, fs.SameAccountRetryCount[100], "账号 100 的计数不应受影响") @@ -351,17 +351,54 @@ func TestHandleFailoverError_SameAccountRetry(t *testing.T) { // 耗尽账号 100 的重试 for i := 0; i < maxSameAccountRetries; i++ { - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) } // 第 maxSameAccountRetries+1 次: 重试耗尽 → 切换 - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) // 再次遇到账号 100,计数仍为 maxSameAccountRetries,条件不满足 → 直接切换 - action = fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action = fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Len(t, mock.calls, 2, "第二次耗尽也应调用 TempUnschedule") }) + + t.Run("尊重账号级retryLimit_配置1次只重试1次", func(t *testing.T) { + // 回归测试:Anthropic 等路径此前硬编码同账号重试 3 次,忽略账号 + // pool_mode_retry_count 配置。此处验证传入 retryLimit=1 时只重试 1 次即切换。 + mock := &mockTempUnscheduler{} + fs := NewFailoverState(5, false) + err := newTestFailoverErr(403, true, false) + const retryLimit = 1 + + // 第 1 次:同账号重试 + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", retryLimit, err) + require.Equal(t, FailoverContinue, action) + require.Equal(t, 1, fs.SameAccountRetryCount[100]) + require.Equal(t, 0, fs.SwitchCount, "首次重试不应切换账号") + require.Empty(t, mock.calls, "未耗尽前不应 TempUnschedule") + + // 第 2 次:已达上限 1 → 不再同账号重试,直接切换 + TempUnschedule + action = fs.HandleFailoverError(context.Background(), mock, 100, "openai", retryLimit, err) + require.Equal(t, FailoverContinue, action) + require.Equal(t, 1, fs.SameAccountRetryCount[100], "重试计数不应超过 retryLimit") + require.Equal(t, 1, fs.SwitchCount, "重试耗尽应切换账号") + require.Contains(t, fs.FailedAccountIDs, int64(100)) + require.Len(t, mock.calls, 1, "重试耗尽应触发 TempUnschedule") + }) + + t.Run("retryLimit为0时立即切换不重试", func(t *testing.T) { + // pool_mode_retry_count=0 表示关闭同账号重试(如 GPT Image 账号)。 + mock := &mockTempUnscheduler{} + fs := NewFailoverState(5, false) + err := newTestFailoverErr(403, true, false) + + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", 0, err) + require.Equal(t, FailoverContinue, action) + require.Equal(t, 0, fs.SameAccountRetryCount[100], "retryLimit=0 不应发生同账号重试") + require.Equal(t, 1, fs.SwitchCount, "应立即切换账号") + require.Len(t, mock.calls, 1, "应立即 TempUnschedule") + }) } // --------------------------------------------------------------------------- @@ -374,7 +411,7 @@ func TestHandleFailoverError_TempUnschedule(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, false, false) // RetryableOnSameAccount=false - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Empty(t, mock.calls) }) @@ -384,10 +421,10 @@ func TestHandleFailoverError_TempUnschedule(t *testing.T) { err := newTestFailoverErr(502, true, false) for i := 0; i < maxSameAccountRetries; i++ { - fs.HandleFailoverError(context.Background(), mock, 42, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 42, "openai", maxSameAccountRetries, err) } // 再次触发时才会执行 TempUnschedule + 切换 - fs.HandleFailoverError(context.Background(), mock, 42, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 42, "openai", maxSameAccountRetries, err) require.Len(t, mock.calls, 1) require.Equal(t, int64(42), mock.calls[0].accountID) @@ -410,7 +447,7 @@ func TestHandleFailoverError_ContextCanceled(t *testing.T) { cancel() // 立即取消 start := time.Now() - action := fs.HandleFailoverError(ctx, mock, 100, "openai", err) + action := fs.HandleFailoverError(ctx, mock, 100, "openai", maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverCanceled, action) @@ -429,7 +466,7 @@ func TestHandleFailoverError_ContextCanceled(t *testing.T) { cancel() // 立即取消 start := time.Now() - action := fs.HandleFailoverError(ctx, mock, 100, service.PlatformAntigravity, err) + action := fs.HandleFailoverError(ctx, mock, 100, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverCanceled, action) @@ -446,10 +483,10 @@ func TestHandleFailoverError_FailedAccountIDs(t *testing.T) { mock := &mockTempUnscheduler{} fs := NewFailoverState(3, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", newTestFailoverErr(500, false, false)) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, newTestFailoverErr(500, false, false)) require.Contains(t, fs.FailedAccountIDs, int64(100)) - fs.HandleFailoverError(context.Background(), mock, 200, "openai", newTestFailoverErr(502, false, false)) + fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, newTestFailoverErr(502, false, false)) require.Contains(t, fs.FailedAccountIDs, int64(200)) require.Len(t, fs.FailedAccountIDs, 2) }) @@ -458,7 +495,7 @@ func TestHandleFailoverError_FailedAccountIDs(t *testing.T) { mock := &mockTempUnscheduler{} fs := NewFailoverState(0, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", newTestFailoverErr(500, false, false)) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, newTestFailoverErr(500, false, false)) require.Equal(t, FailoverExhausted, action) require.Contains(t, fs.FailedAccountIDs, int64(100)) }) @@ -467,7 +504,7 @@ func TestHandleFailoverError_FailedAccountIDs(t *testing.T) { mock := &mockTempUnscheduler{} fs := NewFailoverState(3, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", newTestFailoverErr(400, true, false)) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, newTestFailoverErr(400, true, false)) require.Equal(t, FailoverContinue, action) require.NotContains(t, fs.FailedAccountIDs, int64(100)) }) @@ -476,8 +513,8 @@ func TestHandleFailoverError_FailedAccountIDs(t *testing.T) { mock := &mockTempUnscheduler{} fs := NewFailoverState(5, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", newTestFailoverErr(500, false, false)) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", newTestFailoverErr(500, false, false)) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, newTestFailoverErr(500, false, false)) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, newTestFailoverErr(500, false, false)) require.Len(t, fs.FailedAccountIDs, 1, "map 天然去重") }) } @@ -492,11 +529,11 @@ func TestHandleFailoverError_LastFailoverErr(t *testing.T) { fs := NewFailoverState(3, false) err1 := newTestFailoverErr(500, false, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err1) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err1) require.Equal(t, err1, fs.LastFailoverErr) err2 := newTestFailoverErr(502, false, false) - fs.HandleFailoverError(context.Background(), mock, 200, "openai", err2) + fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, err2) require.Equal(t, err2, fs.LastFailoverErr) }) @@ -505,7 +542,7 @@ func TestHandleFailoverError_LastFailoverErr(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(400, true, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, err, fs.LastFailoverErr) }) } @@ -522,30 +559,30 @@ func TestHandleFailoverError_IntegrationScenario(t *testing.T) { // 1. 账号 100 遇到可重试错误,同账号重试 maxSameAccountRetries 次 retryErr := newTestFailoverErr(400, true, false) for i := 0; i < maxSameAccountRetries; i++ { - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", retryErr) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, retryErr) require.Equal(t, FailoverContinue, action) } require.True(t, fs.ForceCacheBilling, "hasBoundSession=true 应设置 ForceCacheBilling") // 2. 账号 100 超过重试上限 → TempUnschedule + 切换 - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", retryErr) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, retryErr) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SwitchCount) require.Len(t, mock.calls, 1) // 3. 账号 200 遇到不可重试错误 → 直接切换 switchErr := newTestFailoverErr(500, false, false) - action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", switchErr) + action = fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, switchErr) require.Equal(t, FailoverContinue, action) require.Equal(t, 2, fs.SwitchCount) // 4. 账号 300 遇到不可重试错误 → 再切换 - action = fs.HandleFailoverError(context.Background(), mock, 300, "openai", switchErr) + action = fs.HandleFailoverError(context.Background(), mock, 300, "openai", maxSameAccountRetries, switchErr) require.Equal(t, FailoverContinue, action) require.Equal(t, 3, fs.SwitchCount) // 5. 账号 400 → 已耗尽 (SwitchCount=3 >= MaxSwitches=3) - action = fs.HandleFailoverError(context.Background(), mock, 400, "openai", switchErr) + action = fs.HandleFailoverError(context.Background(), mock, 400, "openai", maxSameAccountRetries, switchErr) require.Equal(t, FailoverExhausted, action) // 最终状态验证 @@ -563,21 +600,21 @@ func TestHandleFailoverError_IntegrationScenario(t *testing.T) { // 第一次切换:delay = 0s start := time.Now() - action := fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformAntigravity, err) + action := fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverContinue, action) require.Less(t, elapsed, 200*time.Millisecond, "第一次切换延迟为 0") // 第二次切换:delay = 1s start = time.Now() - action = fs.HandleFailoverError(context.Background(), mock, 200, service.PlatformAntigravity, err) + action = fs.HandleFailoverError(context.Background(), mock, 200, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed = time.Since(start) require.Equal(t, FailoverContinue, action) require.GreaterOrEqual(t, elapsed, 800*time.Millisecond, "第二次切换延迟约 1s") // 第三次:耗尽(无延迟,因为在检查延迟之前就返回了) start = time.Now() - action = fs.HandleFailoverError(context.Background(), mock, 300, service.PlatformAntigravity, err) + action = fs.HandleFailoverError(context.Background(), mock, 300, service.PlatformAntigravity, maxSameAccountRetries, err) elapsed = time.Since(start) require.Equal(t, FailoverExhausted, action) require.Less(t, elapsed, 200*time.Millisecond, "耗尽时不应有延迟") @@ -589,17 +626,17 @@ func TestHandleFailoverError_IntegrationScenario(t *testing.T) { // 第一次:ForceCacheBilling=false err1 := newTestFailoverErr(500, false, false) - fs.HandleFailoverError(context.Background(), mock, 100, "openai", err1) + fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err1) require.False(t, fs.ForceCacheBilling) // 第二次:ForceCacheBilling=true(Antigravity 粘性会话切换) err2 := newTestFailoverErr(500, false, true) - fs.HandleFailoverError(context.Background(), mock, 200, "openai", err2) + fs.HandleFailoverError(context.Background(), mock, 200, "openai", maxSameAccountRetries, err2) require.True(t, fs.ForceCacheBilling, "错误标志应触发 ForceCacheBilling") // 第三次:ForceCacheBilling=false,但状态仍保持 true err3 := newTestFailoverErr(500, false, false) - fs.HandleFailoverError(context.Background(), mock, 300, "openai", err3) + fs.HandleFailoverError(context.Background(), mock, 300, "openai", maxSameAccountRetries, err3) require.True(t, fs.ForceCacheBilling, "不应重置") }) } @@ -614,7 +651,7 @@ func TestHandleFailoverError_EdgeCases(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(0, false, false) - action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) }) @@ -623,7 +660,7 @@ func TestHandleFailoverError_EdgeCases(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, true, false) - action := fs.HandleFailoverError(context.Background(), mock, 0, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, 0, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SameAccountRetryCount[0]) }) @@ -633,7 +670,7 @@ func TestHandleFailoverError_EdgeCases(t *testing.T) { fs := NewFailoverState(3, false) err := newTestFailoverErr(500, true, false) - action := fs.HandleFailoverError(context.Background(), mock, -1, "openai", err) + action := fs.HandleFailoverError(context.Background(), mock, -1, "openai", maxSameAccountRetries, err) require.Equal(t, FailoverContinue, action) require.Equal(t, 1, fs.SameAccountRetryCount[-1]) }) @@ -645,7 +682,7 @@ func TestHandleFailoverError_EdgeCases(t *testing.T) { err := newTestFailoverErr(500, false, false) start := time.Now() - action := fs.HandleFailoverError(context.Background(), mock, 100, "", err) + action := fs.HandleFailoverError(context.Background(), mock, 100, "", maxSameAccountRetries, err) elapsed := time.Since(start) require.Equal(t, FailoverContinue, action) diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 116346b4a6..fb07b098d1 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -448,7 +448,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.handleFailoverExhausted(c, failoverErr, service.PlatformGemini, true) return } - action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, failoverErr) + action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) switch action { case FailoverContinue: continue @@ -868,7 +868,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.handleFailoverExhausted(c, failoverErr, account.Platform, true) return } - action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, failoverErr) + action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) switch action { case FailoverContinue: continue @@ -1454,13 +1454,14 @@ func (h *GatewayHandler) usageUnrestricted(c *gin.Context, ctx context.Context, remaining := h.calculateSubscriptionRemaining(apiKey.Group, subscription) resp["remaining"] = remaining resp["subscription"] = gin.H{ - "daily_usage_usd": subscription.DailyUsageUSD, - "weekly_usage_usd": subscription.WeeklyUsageUSD, - "monthly_usage_usd": subscription.MonthlyUsageUSD, - "daily_limit_usd": apiKey.Group.DailyLimitUSD, - "weekly_limit_usd": apiKey.Group.WeeklyLimitUSD, - "monthly_limit_usd": apiKey.Group.MonthlyLimitUSD, - "expires_at": subscription.ExpiresAt, + "daily_usage_usd": subscription.DailyUsageUSD, + "weekly_usage_usd": subscription.WeeklyUsageUSD, + "monthly_usage_usd": subscription.MonthlyUsageUSD, + "daily_limit_usd": apiKey.Group.DailyLimitUSD, + "weekly_limit_usd": apiKey.Group.WeeklyLimitUSD, + "monthly_limit_usd": apiKey.Group.MonthlyLimitUSD, + "weekly_window_start": subscription.WeeklyWindowStart, + "expires_at": subscription.ExpiresAt, } } diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index f3805f3a53..af9bcdb344 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -254,7 +254,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { h.handleCCFailoverExhausted(c, failoverErr, true) return } - action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, failoverErr) + action := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) switch action { case FailoverContinue: continue diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 5b49ca69a2..8a88d5fa57 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -233,7 +233,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { h.handleResponsesFailoverExhausted(c, failoverErr, true) return } - action := fs.HandleFailoverError(requestCtx, h.gatewayService, account.ID, account.Platform, failoverErr) + action := fs.HandleFailoverError(requestCtx, h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) switch action { case FailoverContinue: continue diff --git a/backend/internal/handler/gateway_handler_usage_test.go b/backend/internal/handler/gateway_handler_usage_test.go new file mode 100644 index 0000000000..b6b0e0efe6 --- /dev/null +++ b/backend/internal/handler/gateway_handler_usage_test.go @@ -0,0 +1,51 @@ +package handler + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "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" +) + +func TestUsageUnrestrictedIncludesWeeklyWindowStart(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/usage", nil) + + weeklyWindowStart := time.Date(2026, time.July, 13, 0, 30, 0, 0, time.FixedZone("UTC+8", 8*60*60)) + c.Set(string(middleware.ContextKeySubscription), &service.UserSubscription{ + WeeklyWindowStart: &weeklyWindowStart, + }) + + handler := &GatewayHandler{} + handler.usageUnrestricted( + c, + context.Background(), + &service.APIKey{Group: &service.Group{ + Name: "Weekly plan", + SubscriptionType: service.SubscriptionTypeSubscription, + }}, + middleware.AuthSubject{}, + nil, + nil, + nil, + ) + + require.Equal(t, http.StatusOK, recorder.Code) + var response struct { + Subscription struct { + WeeklyWindowStart *time.Time `json:"weekly_window_start"` + } `json:"subscription"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + require.NotNil(t, response.Subscription.WeeklyWindowStart) + require.True(t, weeklyWindowStart.Equal(*response.Subscription.WeeklyWindowStart)) +} diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 48110da93f..8489b17f6d 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -220,6 +220,15 @@ func (h *ConcurrencyHelper) TryAcquireUserSlotForAPIKey(ctx context.Context, use return h.withAPIKeySlot(ctx, apiKeyID, releaseFunc), true, nil } +// AcquireOpenAIWSIngressLease bounds the whole client WebSocket lifecycle, +// independently from per-turn user and account slots. +func (h *ConcurrencyHelper) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int) (*service.OpenAIWSIngressLease, bool, error) { + if h == nil || h.concurrencyService == nil { + return nil, false, fmt.Errorf("concurrency service is unavailable") + } + return h.concurrencyService.AcquireOpenAIWSIngressLease(ctx, apiKeyID, maxConnections) +} + // TryAcquireAccountSlot 尝试立即获取账号并发槽位。 // 返回值: (releaseFunc, acquired, error) func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (func(), bool, error) { diff --git a/backend/internal/handler/gateway_helper_fastpath_test.go b/backend/internal/handler/gateway_helper_fastpath_test.go index fecb9b071d..7ae9cb513e 100644 --- a/backend/internal/handler/gateway_helper_fastpath_test.go +++ b/backend/internal/handler/gateway_helper_fastpath_test.go @@ -11,10 +11,13 @@ import ( ) type concurrencyCacheMock struct { - acquireUserSlotFn func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) - acquireAccountSlotFn func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) - releaseUserCalled int32 - releaseAccountCalled int32 + acquireUserSlotFn func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) + acquireAccountSlotFn func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) + acquireIngressLeaseFn func(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) + releaseIngressLeaseFn func(ctx context.Context, apiKeyID int64, leaseID string) error + releaseUserCalled int32 + releaseAccountCalled int32 + releaseIngressCalled int32 } func (m *concurrencyCacheMock) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { @@ -97,6 +100,25 @@ func (m *concurrencyCacheMock) CleanupStaleProcessSlots(ctx context.Context, act return nil } +func (m *concurrencyCacheMock) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) { + if m.acquireIngressLeaseFn != nil { + return m.acquireIngressLeaseFn(ctx, apiKeyID, maxConnections, leaseID) + } + return false, nil +} + +func (m *concurrencyCacheMock) RefreshOpenAIWSIngressLease(context.Context, int64, string) (bool, error) { + return true, nil +} + +func (m *concurrencyCacheMock) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error { + atomic.AddInt32(&m.releaseIngressCalled, 1) + if m.releaseIngressLeaseFn != nil { + return m.releaseIngressLeaseFn(ctx, apiKeyID, leaseID) + } + return nil +} + func TestConcurrencyHelper_TryAcquireUserSlot(t *testing.T) { cache := &concurrencyCacheMock{ acquireUserSlotFn: func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) { diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 86be1062e9..b1653c6b3f 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -482,7 +482,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { if err != nil { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { - failoverAction := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, failoverErr) + failoverAction := fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) switch failoverAction { case FailoverContinue: continue diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 4fd1411b23..b7092fcc42 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -31,6 +31,16 @@ func (h *OpenAIGatewayHandler) GrokVideoGeneration(c *gin.Context) { h.handleGrokMedia(c, service.GrokMediaEndpointVideosGenerations, "") } +// GrokVideoEdit handles asynchronous xAI video edits through Grok groups. +func (h *OpenAIGatewayHandler) GrokVideoEdit(c *gin.Context) { + h.handleGrokMedia(c, service.GrokMediaEndpointVideosEdits, "") +} + +// GrokVideoExtension handles asynchronous xAI video extensions through Grok groups. +func (h *OpenAIGatewayHandler) GrokVideoExtension(c *gin.Context) { + h.handleGrokMedia(c, service.GrokMediaEndpointVideosExtensions, "") +} + // GrokVideoStatus handles xAI video status retrieval through Grok groups. func (h *OpenAIGatewayHandler) GrokVideoStatus(c *gin.Context) { h.handleGrokMedia(c, service.GrokMediaEndpointVideoStatus, c.Param("request_id")) @@ -298,7 +308,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil) - if endpoint == service.GrokMediaEndpointVideosGenerations && strings.TrimSpace(result.ResponseID) != "" { + if endpoint.IsGenerationRequest() && strings.TrimSpace(result.ResponseID) != "" { if err := h.gatewayService.BindGrokMediaVideoRequestAccount(requestCtx, apiKey.GroupID, result.ResponseID, account.ID); err != nil { reqLog.Warn("grok_media.bind_video_request_account_failed", zap.Int64("account_id", account.ID), diff --git a/backend/internal/handler/no_account_error.go b/backend/internal/handler/no_account_error.go index a3bf3b049e..001cef611d 100644 --- a/backend/internal/handler/no_account_error.go +++ b/backend/internal/handler/no_account_error.go @@ -107,3 +107,31 @@ func classifyNoAccountErrorFromGin( } return classifyNoAccountError(ctx, diag, apiKey, routingModel, displayModel, platform) } + +func classifyOpenAICompatibleNoAccountErrorFromGin( + c *gin.Context, + diag service.ModelAvailabilityDiagnoser, + apiKey *service.APIKey, + routingModel string, + displayModel string, +) noAccountErrorClassification { + return classifyNoAccountErrorFromGin( + c, + diag, + apiKey, + routingModel, + displayModel, + openAICompatibleRequestPlatform(apiKey), + ) +} + +func openAICompatibleSelectionErrorForLog(err error, platform string) error { + if err == nil || platform != service.PlatformGrok { + return err + } + message := strings.ReplaceAll(err.Error(), "OpenAI accounts", "Grok accounts") + if message == err.Error() { + return err + } + return fmt.Errorf("%s", message) +} diff --git a/backend/internal/handler/no_account_error_test.go b/backend/internal/handler/no_account_error_test.go index cfe41bb34f..174da82cc7 100644 --- a/backend/internal/handler/no_account_error_test.go +++ b/backend/internal/handler/no_account_error_test.go @@ -4,6 +4,7 @@ package handler import ( "context" + "fmt" "net/http" "net/http/httptest" "testing" @@ -114,6 +115,33 @@ func TestClassifyNoAccountError_ModelNotSupported_Returns404(t *testing.T) { require.Equal(t, int64(42), *fd.calls[0].GroupID) } +func TestClassifyOpenAICompatibleNoAccountError_GrokUsesGrokPlatform(t *testing.T) { + c := newTestGinContextWithRequest() + fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: false}} + groupID := int64(43) + apiKey := &service.APIKey{ + GroupID: &groupID, + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformGrok, + }, + } + + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, fd, apiKey, "grok-4.5", "grok-4.5") + + require.Equal(t, http.StatusNotFound, cls.Status) + require.Equal(t, "model_not_found", cls.ErrType) + require.True(t, cls.ModelNotFound) + require.Len(t, fd.calls, 1) + require.Equal(t, service.PlatformGrok, fd.calls[0].Platform) + + logErr := openAICompatibleSelectionErrorForLog( + fmt.Errorf("no available OpenAI accounts supporting model: grok-4.5"), + service.PlatformGrok, + ) + require.EqualError(t, logErr, "no available Grok accounts supporting model: grok-4.5") +} + func TestClassifyNoAccountError_HasModelSupport_KeepsRoutingMessageGenerationToCaller(t *testing.T) { c := newTestGinContextWithRequest() fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: true}} diff --git a/backend/internal/handler/openai_alpha_search.go b/backend/internal/handler/openai_alpha_search.go new file mode 100644 index 0000000000..a808145235 --- /dev/null +++ b/backend/internal/handler/openai_alpha_search.go @@ -0,0 +1,250 @@ +package handler + +import ( + "context" + "errors" + "net/http" + "strconv" + "strings" + "time" + + pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil" + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +// AlphaSearch proxies the standalone search endpoint used by Codex Responses Lite. +func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { + streamStarted := false + defer h.recoverResponsesPanic(c, &streamStarted) + setOpenAIClientTransportHTTP(c) + requestStart := time.Now() + + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey.Group == nil { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + if apiKey.Group.Platform != service.PlatformOpenAI { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI groups") + return + } + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + reqLog := requestLogger( + c, + "handler.openai_gateway.alpha_search", + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) + if !h.ensureResponsesDependencies(c, reqLog) { + return + } + + body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request) + if err != nil { + if maxErr, ok := extractMaxBytesError(err); ok { + h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) + return + } + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return + } + if len(body) == 0 { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty") + return + } + if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + return + } + + modelResult := gjson.GetBytes(body, "model") + if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return + } + requestedModel := strings.TrimSpace(modelResult.String()) + reqLog = reqLog.With(zap.String("model", requestedModel)) + setOpsRequestContext(c, requestedModel, false) + setOpsEndpointContext(c, "", int16(service.RequestTypeSync)) + + channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestedModel) + forwardBody := openAIModelMappedBody(body, channelMapping.Mapped, channelMapping.MappedModel, h.gatewayService.ReplaceModelInBody) + subscription, _ := middleware2.GetSubscriptionFromContext(c) + service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) + + userRelease, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog) + if !acquired { + return + } + if userRelease != nil { + defer userRelease() + } + + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + searchID := strings.TrimSpace(gjson.GetBytes(body, "id").String()) + sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID) + failedAccountIDs := make(map[int64]struct{}) + var lastFailoverErr *service.UpstreamFailoverError + switchCount := 0 + routingStart := time.Now() + + for { + selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( + c.Request.Context(), + apiKey.GroupID, + "", + sessionHash, + requestedModel, + failedAccountIDs, + service.OpenAIUpstreamTransportHTTPSSE, + service.OpenAIEndpointCapabilityChatCompletions, + false, + false, + service.PlatformOpenAI, + ) + if err != nil || selection == nil || selection.Account == nil { + if len(failedAccountIDs) == 0 { + cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestedModel, requestedModel, service.PlatformOpenAI) + if !cls.ModelNotFound { + markOpsRoutingCapacityLimitedIfNoAvailable(c, err) + } + h.errorResponse(c, cls.Status, cls.ErrType, cls.Message) + return + } + if lastFailoverErr != nil { + h.handleFailoverExhausted(c, lastFailoverErr, false) + } else { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + } + return + } + + account := selection.Account + setOpsSelectedAccount(c, account.ID, account.Platform) + accountRelease, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) + if !acquired { + return + } + service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) + writerSizeBeforeForward := c.Writer.Size() + forwardStart := time.Now() + var result *service.OpenAIForwardResult + result, err = func() (*service.OpenAIForwardResult, error) { + if accountRelease != nil { + defer accountRelease() + } + return h.gatewayService.ForwardAlphaSearch(c.Request.Context(), c, account, forwardBody) + }() + service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds()) + + if err == nil { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil) + if result != nil { + h.recordAlphaSearchUsage(c, apiKey, account, subscription, channelMapping, requestedModel, body, result, subject.UserID) + } + return + } + + var failoverErr *service.UpstreamFailoverError + if !errors.As(err, &failoverErr) { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if c.Writer.Size() == writerSizeBeforeForward { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + } + reqLog.Warn("openai_alpha_search.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + return + } + + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if c.Writer.Size() != writerSizeBeforeForward { + h.handleFailoverExhausted(c, failoverErr, true) + return + } + h.gatewayService.RecordOpenAIAccountSwitch() + failedAccountIDs[account.ID] = struct{}{} + lastFailoverErr = failoverErr + if switchCount >= h.maxAccountSwitches { + h.handleFailoverExhausted(c, failoverErr, false) + return + } + switchCount++ + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + h.handleFailoverExhausted(c, failoverErr, false) + return + } + reqLog.Warn("openai_alpha_search.upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", switchCount), + ) + } +} + +// recordAlphaSearchUsage 为一次成功的 alpha/search 网页搜索落按次计费用量行 +// (上游不返回 usage 字段,按 WebSearchCalls 走分组单价 × 倍率的按次口径)。 +// 与 images 一致使用 mandatory 池提交,池满时同步兜底执行,保证扣费不丢。 +func (h *OpenAIGatewayHandler) recordAlphaSearchUsage( + c *gin.Context, + apiKey *service.APIKey, + account *service.Account, + subscription *service.UserSubscription, + channelMapping service.ChannelMappingResult, + requestedModel string, + body []byte, + result *service.OpenAIForwardResult, + userID int64, +) { + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + + h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ + Result: result, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + ChannelUsageFields: channelMapping.ToUsageFields(requestedModel, result.UpstreamModel), + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.alpha_search"), + zap.Int64("user_id", userID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("model", requestedModel), + zap.Int64("account_id", account.ID), + ).Error("openai_alpha_search.record_usage_failed", zap.Error(err)) + } + }) +} diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index f5f2522e49..636e143740 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -5,6 +5,7 @@ import ( "errors" "net/http" "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/ip" @@ -150,11 +151,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { ) if err != nil { reqLog.Warn("openai_chat_completions.account_select_failed", - zap.Error(err), + zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), ) if len(failedAccountIDs) == 0 { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimitedIfNoAvailable(c, err) } @@ -170,7 +171,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { } } if selection == nil || selection.Account == nil { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) } @@ -298,7 +299,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { userAgent := c.GetHeader("User-Agent") clientIP := ip.GetClientIP(c) inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) cyberBlocked := service.GetOpsCyberPolicy(c) != nil @@ -337,14 +338,22 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { } // resolveOpenAIUpstreamEndpoint returns the actual upstream endpoint for an -// OpenAI account, used by every OpenAI usage-recording site. APIKey accounts -// whose upstream is forced or probed to not support the Responses API are -// served directly via /v1/chat/completions (the raw chat path) regardless of -// the inbound endpoint; everything else goes through the Responses API. -func resolveOpenAIUpstreamEndpoint(c *gin.Context, account *service.Account) string { +// OpenAI-compatible account. A forwarding result is authoritative because a +// single inbound route may choose raw Chat or a Responses bridge at runtime. +// The account-based derivation remains as a fallback for existing callers and +// forwarding paths that do not report their endpoint yet. +func resolveOpenAIUpstreamEndpoint(c *gin.Context, account *service.Account, result *service.OpenAIForwardResult) string { + if result != nil { + if endpoint := strings.TrimSpace(result.UpstreamEndpoint); endpoint != "" { + return endpoint + } + } + if endpoint := service.GetActualOpenAIUpstreamEndpoint(c); endpoint != "" { + return endpoint + } if account != nil && account.Type == service.AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) { - return "/v1/chat/completions" + return EndpointChatCompletions } return GetUpstreamEndpoint(c, account.Platform) } diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index e64c555d14..1c1357cbfa 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -15,11 +15,13 @@ import ( // Codex CLI and the Codex desktop app refresh their model picker from // GET {base_url}/models?client_version=... (custom provider mode) or // GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land -// here. The manifest is proxied verbatim from the ChatGPT backend with a -// schedulable OAuth account's credentials, so clients pointed at the gateway -// see the account's real, always-current model entitlements instead of a -// frozen local cache. +// here. The manifest is proxied verbatim from the selected account's ChatGPT +// backend or custom API key upstream. API key manifests use a short-lived, +// asynchronously revalidated cache to tolerate canceled client requests. func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { + if c.Request.Context().Err() != nil { + return + } apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey.Group == nil { h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required") @@ -30,24 +32,54 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { return } - account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "") - if err != nil { - h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") - return + maxAccountSwitches := h.maxAccountSwitches + if maxAccountSwitches <= 0 { + maxAccountSwitches = 3 } + failedAccountIDs := make(map[int64]struct{}) + switchCount := 0 + var lastUpstreamErr error - manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) - if err != nil { - h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) - return - } + for { + account, err := h.gatewayService.SelectAccountForModelWithExclusions(c.Request.Context(), apiKey.GroupID, "", "", failedAccountIDs) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if lastUpstreamErr != nil { + h.errorResponse(c, infraerrors.Code(lastUpstreamErr), "upstream_error", infraerrors.Message(lastUpstreamErr)) + return + } + h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") + return + } - if manifest.ETag != "" { - c.Header("ETag", manifest.ETag) - } - if manifest.NotModified { - c.Status(http.StatusNotModified) + manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if service.IsRetryableCodexModelsManifestError(err) && switchCount < maxAccountSwitches { + failedAccountIDs[account.ID] = struct{}{} + switchCount++ + lastUpstreamErr = err + continue + } + h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) + return + } + if c.Request.Context().Err() != nil { + return + } + + if manifest.ETag != "" { + c.Header("ETag", manifest.ETag) + } + if manifest.NotModified { + c.Status(http.StatusNotModified) + return + } + c.Data(http.StatusOK, "application/json", manifest.Body) return } - c.Data(http.StatusOK, "application/json", manifest.Body) } diff --git a/backend/internal/handler/openai_codex_models_handler_test.go b/backend/internal/handler/openai_codex_models_handler_test.go new file mode 100644 index 0000000000..ba74a5869f --- /dev/null +++ b/backend/internal/handler/openai_codex_models_handler_test.go @@ -0,0 +1,288 @@ +package handler + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +type codexModelsFailoverAccountRepo struct { + service.AccountRepository + accounts []service.Account +} + +func (r codexModelsFailoverAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + for i := range r.accounts { + if r.accounts[i].ID == id { + account := r.accounts[i] + return &account, nil + } + } + return nil, service.ErrNoAvailableAccounts +} + +func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) { + accounts := make([]service.Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform { + accounts = append(accounts, account) + } + } + return accounts, nil +} + +type codexModelsFailoverHTTPUpstream struct { + service.HTTPUpstream + mu sync.Mutex + accountIDs []int64 + firstErr error + firstStatus int + statuses map[int64]int +} + +func (u *codexModelsFailoverHTTPUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + u.mu.Lock() + u.accountIDs = append(u.accountIDs, accountID) + u.mu.Unlock() + + status, hasStatus := u.statuses[accountID] + if accountID == 1 || hasStatus { + if u.firstErr != nil { + return nil, u.firstErr + } + if !hasStatus { + status = u.firstStatus + } + if status == 0 { + status = http.StatusServiceUnavailable + } + return &http.Response{ + StatusCode: status, + Status: http.StatusText(status), + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader( + `{"error":{"message":"No available OpenAI accounts","type":"upstream_error"}}`, + )), + }, nil + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"gpt-5.6-sol"}]}`)), + }, nil +} + +func (u *codexModelsFailoverHTTPUpstream) calls() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.accountIDs...) +} + +func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil).WithContext(ctx) + + h := &OpenAIGatewayHandler{} + h.CodexModels(c) + + if c.Writer.Written() { + t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String()) + } +} + +func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) { + retryableStatuses := []int{ + http.StatusTooManyRequests, + http.StatusInternalServerError, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout, + } + for _, status := range retryableStatuses { + t.Run(http.StatusText(status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want { + t.Fatalf("body: got %q, want %q", got, want) + } + }) + } +} + +func TestCodexModelsFailsOverFromUpstreamTransportError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } +} + +func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) { + statuses := []int{ + http.StatusBadRequest, + http.StatusUnauthorized, + http.StatusForbidden, + http.StatusNotFound, + 600, + } + for _, status := range statuses { + t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + }) + } +} + +func TestCodexModelsDoesNotFailOverFromUpstreamConfigurationError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = errors.New("invalid proxy URL") + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } +} + +func TestCodexModelsReturnsLastUpstreamErrorWhenAccountsAreExhausted(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusGatewayTimeout, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the last upstream error: %s", body) + } +} + +func TestCodexModelsHonorsAccountSwitchLimit(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandlerWithAccountCount(http.StatusServiceUnavailable, 4, 2) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusBadGateway, + 3: http.StatusGatewayTimeout, + 4: http.StatusInternalServerError, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2, 3}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the limit-ending upstream error: %s", body) + } +} + +func newCodexModelsFailoverTestHandler(firstStatus int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + return newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, 2, 3) +} + +func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount, maxSwitches int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + gin.SetMode(gin.TestMode) + groupID := int64(42) + accounts := make([]service.Account, 0, accountCount) + for i := 1; i <= accountCount; i++ { + accounts = append(accounts, service.Account{ + ID: int64(i), + Name: fmt.Sprintf("upstream-%d", i), + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Priority: i - 1, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": fmt.Sprintf("sk-%d", i), + "base_url": fmt.Sprintf("https://upstream-%d.example/v1", i), + }, + }) + } + upstream := &codexModelsFailoverHTTPUpstream{firstStatus: firstStatus} + cfg := &config.Config{RunMode: config.RunModeSimple} + gatewayService := service.NewOpenAIGatewayService( + codexModelsFailoverAccountRepo{accounts: accounts}, + nil, nil, nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + return &OpenAIGatewayHandler{gatewayService: gatewayService, maxAccountSwitches: maxSwitches}, upstream, groupID +} + +func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder { + t.Helper() + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + GroupID: &groupID, + Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}, + }) + + handler.CodexModels(c) + return recorder +} + +func equalInt64Slices(got, want []int64) bool { + if len(got) != len(want) { + return false + } + for i := range got { + if got[i] != want[i] { + return false + } + } + return true +} diff --git a/backend/internal/handler/openai_gateway_compact_body_signal_test.go b/backend/internal/handler/openai_gateway_compact_body_signal_test.go index a4bfb90466..a44d47c856 100644 --- a/backend/internal/handler/openai_gateway_compact_body_signal_test.go +++ b/backend/internal/handler/openai_gateway_compact_body_signal_test.go @@ -23,46 +23,61 @@ func newCompactBodySignalTestContext(t *testing.T, path string, body []byte) *gi return c } -// body-signal 提升后必须与 path-based compact 走同一条链路: -// path 改写、requireCompact 判定、stream/store/prompt_cache_key 归一化删除。 -// 回归防护:若 stream 字段存活,Forward 会用流式 handler 解析 compact 的 -// JSON 响应,导致 "stream ended before a terminal event" 的换号 failover 风暴。 -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalPromoted(t *testing.T) { +func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2StaysOnResponses(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{ - "model":"gpt-5.5", + "model":"gpt-5.6-sol", "stream":true, "store":true, "prompt_cache_key":"pck-signal-1", + "reasoning":{"effort":"max","context":"all_turns"}, "input":[ {"type":"message","role":"user","content":"hello"}, {"type":"compaction_trigger"} ] }`) c := newCompactBodySignalTestContext(t, "/v1/responses", body) + c.Request.Header.Set("x-codex-beta-features", "responses_websockets_v2, remote_compaction_v2, another_feature") normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) require.True(t, ok) - require.Equal(t, "/v1/responses/compact", c.Request.URL.Path) - require.True(t, isOpenAIRemoteCompactPath(c)) - - require.False(t, gjson.GetBytes(normalized, "stream").Exists()) - require.False(t, gjson.GetBytes(normalized, "store").Exists()) - require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists()) - require.Equal(t, "gpt-5.5", gjson.GetBytes(normalized, "model").String()) - require.True(t, gjson.GetBytes(normalized, "input").IsArray()) + require.Equal(t, "/v1/responses", c.Request.URL.Path) + require.False(t, isOpenAIRemoteCompactPath(c)) + require.Equal(t, body, normalized) + require.True(t, gjson.GetBytes(normalized, "stream").Bool()) + require.True(t, gjson.GetBytes(normalized, "store").Bool()) + require.Equal(t, "pck-signal-1", gjson.GetBytes(normalized, "prompt_cache_key").String()) + require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(normalized, "reasoning.context").String()) reqStream, streamOK := parseOpenAICompatibleStream(normalized) require.True(t, streamOK) - require.False(t, reqStream) + require.True(t, reqStream) - seed, exists := c.Get(service.OpenAICompactSessionSeedKeyForTest()) - require.True(t, exists) - require.Equal(t, "pck-signal-1", seed) + _, seedExists := c.Get(service.OpenAICompactSessionSeedKeyForTest()) + require.False(t, seedExists) + _, streamMarkerExists := c.Get(service.OpenAICompactClientStreamKeyForTest()) + require.False(t, streamMarkerExists) } -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlash(t *testing.T) { +func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2PathAliasesStayOnResponses(t *testing.T) { + h := &OpenAIGatewayHandler{} + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"input":[{"type":"compaction_trigger"}]}`) + for _, path := range []string{"/v1/responses/", "/backend-api/codex/responses"} { + t.Run(path, func(t *testing.T) { + c := newCompactBodySignalTestContext(t, path, body) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") + + normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) + require.True(t, ok) + require.Equal(t, path, c.Request.URL.Path) + require.Equal(t, body, normalized) + }) + } +} + +func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlashPromoted(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`) c := newCompactBodySignalTestContext(t, "/v1/responses/", body) @@ -82,6 +97,64 @@ func TestNormalizeOpenAIResponsesCompactRequest_CodexDirectAliasPromoted(t *test require.Equal(t, "/backend-api/codex/responses/compact", c.Request.URL.Path) } +func TestNormalizeOpenAIResponsesCompactRequest_NonRemoteV2BodySignalPromoted(t *testing.T) { + h := &OpenAIGatewayHandler{} + tests := []struct { + name string + body []byte + betaHeader string + wantMarked bool + }{ + { + name: "no_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + wantMarked: true, + }, + { + name: "unrelated_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "responses_websockets_v2", + wantMarked: true, + }, + { + name: "wrong_case_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "REMOTE_COMPACTION_V2", + wantMarked: true, + }, + { + name: "stream_false", + body: []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "remote_compaction_v2", + }, + { + name: "stream_absent", + body: []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`), + betaHeader: "remote_compaction_v2", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := newCompactBodySignalTestContext(t, "/v1/responses", tt.body) + if tt.betaHeader != "" { + c.Request.Header.Set("x-codex-beta-features", tt.betaHeader) + } + + normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), tt.body) + require.True(t, ok) + require.Equal(t, "/v1/responses/compact", c.Request.URL.Path) + require.False(t, gjson.GetBytes(normalized, "stream").Exists()) + + marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) + require.Equal(t, tt.wantMarked, exists) + if tt.wantMarked { + require.Equal(t, true, marked) + } + }) + } +} + func TestNormalizeOpenAIResponsesCompactRequest_NoTriggerUntouched(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"message","role":"user","content":"hello"}]}`) @@ -99,6 +172,7 @@ func TestNormalizeOpenAIResponsesCompactRequest_PathBasedNoDoubleSuffix(t *testi h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","stream":true,"store":true,"input":[{"type":"message","role":"user","content":"hello"}]}`) c := newCompactBodySignalTestContext(t, "/v1/responses/compact", body) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) require.True(t, ok) @@ -118,36 +192,6 @@ func TestNormalizeOpenAIResponsesCompactRequest_SubpathNotPromoted(t *testing.T) require.Equal(t, body, normalized) } -// 回归 #3875:body-signal 原始请求 stream:true 时必须标记 client-stream, -// 供响应写回阶段把上游 unary JSON 合成回 Codex remote compact v2 所需的 SSE。 -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamTrueMarksClientStream(t *testing.T) { - h := &OpenAIGatewayHandler{} - body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`) - c := newCompactBodySignalTestContext(t, "/v1/responses", body) - - _, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) - require.True(t, ok) - - marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) - require.True(t, exists) - require.Equal(t, true, marked) -} - -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamFalseNotMarked(t *testing.T) { - h := &OpenAIGatewayHandler{} - for name, body := range map[string][]byte{ - "stream_false": []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`), - "stream_absent": []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`), - } { - c := newCompactBodySignalTestContext(t, "/v1/responses", body) - _, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) - require.True(t, ok, name) - require.Equal(t, "/v1/responses/compact", c.Request.URL.Path, name) - _, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) - require.False(t, exists, "case %s 不应标记 client-stream", name) - } -} - // path-based compact(Codex v1 unary 协议)即使 body 带 stream:true 也不标记, // 保持 JSON 写回行为不变。 func TestNormalizeOpenAIResponsesCompactRequest_PathBasedStreamTrueNotMarked(t *testing.T) { diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index 0461017067..0a010cc176 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -115,8 +115,9 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) if err != nil { - reqLog.Warn("openai_count_tokens.account_select_failed", zap.Error(err)) - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI) + requestPlatform := openAICompatibleRequestPlatform(apiKey) + reqLog.Warn("openai_count_tokens.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform))) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimitedIfNoAvailable(c, err) } @@ -124,7 +125,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { return } if selection == nil || selection.Account == nil { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 83e644d857..231363a8a3 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -351,7 +351,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { ) if err != nil { reqLog.Warn("openai.account_select_failed", - zap.Error(err), + zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), ) if len(failedAccountIDs) == 0 { @@ -360,7 +360,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "compact_not_supported", "No available OpenAI accounts support /responses/compact", streamStarted) return } - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimitedIfNoAvailable(c, err) } @@ -375,7 +375,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { return } if selection == nil || selection.Account == nil { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) } @@ -522,7 +522,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { clientIP := ip.GetClientIP(c) requestPayloadHash := service.HashUsageRequestPayload(body) inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 @@ -580,21 +580,33 @@ func isBareOpenAIResponsesPath(c *gin.Context) bool { return strings.HasSuffix(normalizedPath, "/responses") } -// normalizeOpenAIResponsesCompactRequest 统一处理两种入站 compact 形态: -// path-based(POST /v1/responses/compact)与 Codex remote compact v2 的 -// body-signal(普通 POST /v1/responses 的 input 中携带 type=compaction_trigger, -// 见 #3777)。body-signal 命中时在 stream 解析、compact body 归一化与 -// requireCompact 调度判定之前改写 URL path,使后续全部链路(含 passthrough -// 分支与上游 URL 构建)与 path-based 完全一致。 +func isOpenAIRemoteCompactionV2Request(c *gin.Context, body []byte) bool { + stream, valid := parseOpenAICompatibleStream(body) + if !valid || !stream || c == nil || c.Request == nil { + return false + } + for _, header := range c.Request.Header.Values("x-codex-beta-features") { + for _, feature := range strings.Split(header, ",") { + if strings.TrimSpace(feature) == "remote_compaction_v2" { + return true + } + } + } + return false +} + +// normalizeOpenAIResponsesCompactRequest keeps Codex remote compaction v2 on +// its native streaming /responses wire and preserves the legacy body-signal +// promotion for clients that do not explicitly advertise that protocol. // 返回归一化后的 body;ok=false 表示错误响应已写出,调用方应直接 return。 func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) { isCompactRequest := service.IsOpenAIResponsesCompactPathForTest(c) if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) { + if isOpenAIRemoteCompactionV2Request(c, body) { + return body, true + } c.Request.URL.Path = strings.TrimRight(c.Request.URL.Path, "/") + "/compact" isCompactRequest = true - // Codex remote compact v2 的原始请求是流式 /responses:白名单归一化会删除 - // stream 并让上游走 unary JSON,但客户端仍按 SSE 消费响应。记录原始 - // stream 意图,响应写回阶段据此把 JSON 合成回 SSE(#3875)。 clientStream := gjson.GetBytes(body, "stream").Bool() if clientStream { service.MarkOpenAICompactClientStream(c) @@ -843,12 +855,12 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { ) if err != nil { reqLog.Warn("openai_messages.account_select_failed", - zap.Error(err), + zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), ) if len(failedAccountIDs) == 0 { if err != nil { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimitedIfNoAvailable(c, err) } @@ -865,7 +877,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } } if selection == nil || selection.Account == nil { - cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI) + cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) } @@ -994,7 +1006,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { clientIP := ip.GetClientIP(c) requestPayloadHash := service.HashUsageRequestPayload(body) inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) cyberBlocked := service.GetOpsCyberPolicy(c) != nil @@ -1287,10 +1299,36 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { wsConn.SetReadLimit(service.ResolveOpenAIWSClientReadLimitBytes(h.cfg)) ctx := c.Request.Context() + maxIngressConnections := 0 + if h.cfg != nil { + maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey + } + ingressLease, ingressLeaseAcquired, ingressLeaseErr := h.concurrencyHelper.AcquireOpenAIWSIngressLease(ctx, apiKey.ID, maxIngressConnections) + if ingressLeaseErr != nil { + reqLog.Error("openai.websocket_ingress_lease_acquire_failed", zap.Error(ingressLeaseErr)) + closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to reserve websocket ingress capacity") + return + } + if !ingressLeaseAcquired { + reqLog.Info("openai.websocket_ingress_capacity_rejected", zap.Int("max_ingress_connections_per_api_key", maxIngressConnections)) + closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "too many open websocket connections, please retry later") + return + } + if ingressLease != nil { + defer ingressLease.Release() + ctx = ingressLease.Context() + c.Request = c.Request.WithContext(ctx) + } + readCtx, cancel := context.WithTimeout(ctx, 30*time.Second) msgType, firstMessage, err := wsConn.Read(readCtx) cancel() if err != nil { + if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) { + reqLog.Warn("openai.websocket_ingress_lease_lost_before_first_message", zap.Error(err)) + closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "websocket ingress capacity lease lost; please reconnect") + return + } closeStatus, closeReason := summarizeWSCloseErrorForLog(err) reqLog.Warn("openai.websocket_read_first_message_failed", zap.Error(err), @@ -1444,7 +1482,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { ) if err != nil { reqLog.Warn("openai.websocket_account_select_failed", - zap.Error(err), + zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), ) if lastFailoverErr != nil { @@ -1601,7 +1639,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs) inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) { @@ -1680,6 +1718,25 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { continue } + if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) { + reqLog.Warn("openai.websocket_ingress_lease_lost", + zap.Int64("account_id", account.ID), + zap.Error(err), + ) + closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "websocket ingress capacity lease lost; please reconnect") + return + } + + var closeErr *service.OpenAIWSClientCloseError + if errors.As(err, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure { + reqLog.Info("openai.websocket_ingress_closed_normally", + zap.Int64("account_id", account.ID), + zap.String("reason", closeErr.Reason()), + ) + closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason()) + return + } + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) closeStatus, closeReason := summarizeWSCloseErrorForLog(err) reqLog.Warn("openai.websocket_proxy_failed", @@ -1688,7 +1745,6 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { zap.String("close_status", closeStatus), zap.String("close_reason", closeReason), ) - var closeErr *service.OpenAIWSClientCloseError if errors.As(err, &closeErr) { closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason()) return @@ -2035,7 +2091,8 @@ func openAIForwardErrorAlreadyCommunicated(c *gin.Context, writerSizeBeforeForwa } // 与快照同口径:排除 compact 心跳字节,避免"仅心跳写出"被误判为 // 响应已写出(#3887)。 - if service.OpenAICompactKeepaliveAdjustedWrittenSize(c) == writerSizeBeforeForward { + if service.OpenAICompactKeepaliveAdjustedWrittenSize(c) == writerSizeBeforeForward || + service.OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) == writerSizeBeforeForward { return false } @@ -2441,7 +2498,7 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey var accountID int64 if account != nil { accountID = account.ID - upstreamEndpoint = resolveOpenAIUpstreamEndpoint(c, account) + upstreamEndpoint = resolveOpenAIUpstreamEndpoint(c, account, nil) } stream := false if v, ok := c.Get(opsStreamKey); ok { diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index b7f43079ef..e4b594c0b8 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -8,6 +8,7 @@ import ( "net/http/httptest" "strings" "sync" + "sync/atomic" "testing" "time" @@ -414,11 +415,13 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { SonnetMappedModel: "gpt-5.2", ExactModelMappings: map[string]string{ "claude-sonnet-4-5-20250929": "gpt-5.4-mini-high", + "claude-fable-5": "gpt-5.6-sol", }, }, }, } require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-fable-5")) }) t.Run("uses_family_default_when_no_override", func(t *testing.T) { @@ -711,6 +714,68 @@ func TestOpenAIResponsesWebSocket_InvalidUpgradeDoesNotSetTransport(t *testing.T require.Equal(t, service.OpenAIClientTransportUnknown, service.GetOpenAIClientTransport(c)) } +func TestOpenAIResponsesWebSocket_IngressCapacityRejected(t *testing.T) { + gin.SetMode(gin.TestMode) + cache := &concurrencyCacheMock{ + acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) { + return false, nil + }, + } + h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache) + h.cfg = &config.Config{} + h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1 + wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1}) + defer wsServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil) + cancelDial() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, _, err = clientConn.Read(readCtx) + cancelRead() + var closeErr coderws.CloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusTryAgainLater, closeErr.Code) +} + +func TestOpenAIResponsesWebSocket_IngressLeaseReleasedOnEarlyReturn(t *testing.T) { + gin.SetMode(gin.TestMode) + cache := &concurrencyCacheMock{ + acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) { + return true, nil + }, + } + h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache) + h.cfg = &config.Config{} + h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1 + wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1}) + defer wsServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil) + cancelDial() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageBinary, []byte("not a response.create frame")) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, _, err = clientConn.Read(readCtx) + cancelRead() + var closeErr coderws.CloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code) + require.Eventually(t, func() bool { + return atomic.LoadInt32(&cache.releaseIngressCalled) == 1 + }, time.Second, 10*time.Millisecond) +} + func TestOpenAIResponsesWebSocket_RejectsMessageIDAsPreviousResponseID(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index 5868f7f35b..c5982fb7d1 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -142,6 +142,9 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + stopJSONKeepalive := func() {} + jsonKeepaliveStarted := false + defer func() { stopJSONKeepalive() }() for { reqLog.Debug("openai.images.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) @@ -210,8 +213,12 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { } service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) + if !parsed.Stream && !jsonKeepaliveStarted { + stopJSONKeepalive = service.StartOpenAIImagesJSONKeepalive(c, h.openAIImagesJSONKeepaliveInterval()) + jsonKeepaliveStarted = true + } forwardStart := time.Now() - writerSizeBeforeForward := c.Writer.Size() + writerSizeBeforeForward := service.OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) result, err := func() (*service.OpenAIForwardResult, error) { defer func() { if accountReleaseFunc != nil { @@ -258,7 +265,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) - if c.Writer.Size() != writerSizeBeforeForward { + if service.OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeForward { reqLog.Warn("openai.images.upstream_failover_skipped_after_flush", zap.Int64("account_id", account.ID), zap.Int("upstream_status", failoverErr.StatusCode), @@ -383,6 +390,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { } } +func (h *OpenAIGatewayHandler) openAIImagesJSONKeepaliveInterval() time.Duration { + if h.cfg == nil || h.cfg.Gateway.ImageNonstreamKeepaliveInterval <= 0 { + return 0 + } + return time.Duration(h.cfg.Gateway.ImageNonstreamKeepaliveInterval) * time.Second +} + func isMultipartImagesContentType(contentType string) bool { return strings.HasPrefix(strings.ToLower(strings.TrimSpace(contentType)), "multipart/form-data") } diff --git a/backend/internal/handler/ops_capture_writer_nil_test.go b/backend/internal/handler/ops_capture_writer_nil_test.go new file mode 100644 index 0000000000..88aa7c043f --- /dev/null +++ b/backend/internal/handler/ops_capture_writer_nil_test.go @@ -0,0 +1,90 @@ +package handler + +import ( + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) { + w := &opsCaptureWriter{} + w.ResponseWriter = nil + + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Status()) + }) + assert.NotPanics(t, func() { + assert.Equal(t, -1, w.Size()) + }) + assert.NotPanics(t, func() { + assert.False(t, w.Written()) + }) + assert.NotPanics(t, func() { + n, err := w.Write([]byte("test")) + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + n, err := w.WriteString("test") + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + h := w.Header() + assert.NotNil(t, h) + }) + assert.NotPanics(t, func() { + w.WriteHeader(200) + }) + assert.NotPanics(t, func() { + w.WriteHeaderNow() + }) + assert.NotPanics(t, func() { + w.Flush() + }) + assert.NotPanics(t, func() { + conn, rw, err := w.Hijack() + assert.Nil(t, conn) + assert.Nil(t, rw) + assert.Error(t, err) + }) + assert.NotPanics(t, func() { + ch := w.CloseNotify() + assert.NotNil(t, ch) + }) + assert.NotPanics(t, func() { + p := w.Pusher() + assert.Nil(t, p) + }) +} + +func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + outerStatus := -1 + router.Use(func(c *gin.Context) { + c.Next() + outerStatus = c.Writer.Status() + }) + router.Use(OpsErrorLoggerMiddleware(nil)) + router.GET("/compact", func(c *gin.Context) { + service.MarkOpenAICompactClientStream(c) + stop := service.StartOpenAICompactSSEKeepalive(c, time.Hour) + defer stop() + c.Status(http.StatusOK) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/compact", nil) + require.NotPanics(t, func() { + router.ServeHTTP(recorder, request) + }) + require.Equal(t, http.StatusOK, outerStatus) + require.Equal(t, http.StatusOK, recorder.Code) +} diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 5a1e57ff7d..64aa3ba495 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -1,11 +1,14 @@ package handler import ( + "bufio" "bytes" "context" "encoding/json" "errors" "log" + "net" + "net/http" "runtime" "runtime/debug" "strconv" @@ -498,7 +501,82 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) { opsCaptureWriterPool.Put(w) } +func (w *opsCaptureWriter) Status() int { + if w.ResponseWriter == nil { + return 0 + } + return w.ResponseWriter.Status() +} + +func (w *opsCaptureWriter) Size() int { + if w.ResponseWriter == nil { + return -1 + } + return w.ResponseWriter.Size() +} + +func (w *opsCaptureWriter) Written() bool { + if w.ResponseWriter == nil { + return false + } + return w.ResponseWriter.Written() +} + +func (w *opsCaptureWriter) Header() http.Header { + if w.ResponseWriter == nil { + return http.Header{} + } + return w.ResponseWriter.Header() +} + +func (w *opsCaptureWriter) WriteHeader(code int) { + if w.ResponseWriter == nil { + return + } + w.ResponseWriter.WriteHeader(code) +} + +func (w *opsCaptureWriter) WriteHeaderNow() { + if w.ResponseWriter == nil { + return + } + w.ResponseWriter.WriteHeaderNow() +} + +func (w *opsCaptureWriter) Flush() { + if w.ResponseWriter == nil { + return + } + w.ResponseWriter.Flush() +} + +func (w *opsCaptureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + if w.ResponseWriter == nil { + return nil, nil, errors.New("response writer released") + } + return w.ResponseWriter.Hijack() +} + +func (w *opsCaptureWriter) CloseNotify() <-chan bool { + if w.ResponseWriter == nil { + ch := make(chan bool) + close(ch) + return ch + } + return w.ResponseWriter.CloseNotify() +} + +func (w *opsCaptureWriter) Pusher() http.Pusher { + if w.ResponseWriter == nil { + return nil + } + return w.ResponseWriter.Pusher() +} + func (w *opsCaptureWriter) Write(b []byte) (int, error) { + if w.ResponseWriter == nil { + return 0, nil + } if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit { remaining := w.limit - w.buf.Len() if len(b) > remaining { @@ -511,6 +589,9 @@ func (w *opsCaptureWriter) Write(b []byte) (int, error) { } func (w *opsCaptureWriter) WriteString(s string) (int, error) { + if w.ResponseWriter == nil { + return 0, nil + } if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit { remaining := w.limit - w.buf.Len() if len(s) > remaining { diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index a267d73724..1ad054da75 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -9,7 +9,6 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/internal/payment" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" - "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/pkg/response" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" @@ -19,15 +18,13 @@ import ( // PaymentHandler handles user-facing payment requests. type PaymentHandler struct { - channelService *service.ChannelService paymentService *service.PaymentService configService *service.PaymentConfigService } // NewPaymentHandler creates a new PaymentHandler. -func NewPaymentHandler(paymentService *service.PaymentService, configService *service.PaymentConfigService, channelService *service.ChannelService) *PaymentHandler { +func NewPaymentHandler(paymentService *service.PaymentService, configService *service.PaymentConfigService) *PaymentHandler { return &PaymentHandler{ - channelService: channelService, paymentService: paymentService, configService: configService, } @@ -91,17 +88,6 @@ func (h *PaymentHandler) GetPlans(c *gin.Context) { response.Success(c, result) } -// GetChannels returns enabled payment channels. -// GET /api/v1/payment/channels -func (h *PaymentHandler) GetChannels(c *gin.Context) { - channels, _, err := h.channelService.List(c.Request.Context(), pagination.PaginationParams{Page: 1, PageSize: 1000}, "active", "") - if err != nil { - response.ErrorFrom(c, err) - return - } - response.Success(c, channels) -} - // GetCheckoutInfo returns all data the payment page needs in a single call: // payment methods with limits, subscription plans, and configuration. // GET /api/v1/payment/checkout-info diff --git a/backend/internal/handler/payment_handler_resume_test.go b/backend/internal/handler/payment_handler_resume_test.go index 21fd8ad763..c902f390e0 100644 --- a/backend/internal/handler/payment_handler_resume_test.go +++ b/backend/internal/handler/payment_handler_resume_test.go @@ -119,7 +119,7 @@ func TestVerifyOrderPublicReturnsLegacyOrderState(t *testing.T) { require.NoError(t, err) paymentSvc := service.NewPaymentService(client, payment.NewRegistry(), nil, nil, nil, nil, nil, nil, nil) - h := NewPaymentHandler(paymentSvc, nil, nil) + h := NewPaymentHandler(paymentSvc, nil) recorder := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(recorder) @@ -219,7 +219,7 @@ func TestResolveOrderPublicByResumeTokenReturnsFrontendContractFields(t *testing configSvc := service.NewPaymentConfigService(client, nil, []byte("0123456789abcdef0123456789abcdef")) paymentSvc := service.NewPaymentService(client, payment.NewRegistry(), nil, nil, nil, configSvc, nil, nil, nil) - h := NewPaymentHandler(paymentSvc, nil, nil) + h := NewPaymentHandler(paymentSvc, nil) recorder := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(recorder) @@ -307,7 +307,7 @@ func TestResolveOrderPublicByResumeTokenReturnsBadRequestForMismatchedToken(t *t configSvc := service.NewPaymentConfigService(client, nil, []byte("0123456789abcdef0123456789abcdef")) paymentSvc := service.NewPaymentService(client, payment.NewRegistry(), nil, nil, nil, configSvc, nil, nil, nil) - h := NewPaymentHandler(paymentSvc, nil, nil) + h := NewPaymentHandler(paymentSvc, nil) recorder := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(recorder) @@ -347,7 +347,7 @@ func TestVerifyOrderPublicRejectsBlankOutTradeNo(t *testing.T) { t.Cleanup(func() { _ = client.Close() }) paymentSvc := service.NewPaymentService(client, payment.NewRegistry(), nil, nil, nil, nil, nil, nil, nil) - h := NewPaymentHandler(paymentSvc, nil, nil) + h := NewPaymentHandler(paymentSvc, nil) recorder := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(recorder) diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index cfbb72554c..67380fed33 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -164,7 +164,7 @@ var ProviderSet = wire.NewSet( admin.NewDashboardHandler, admin.NewUserHandler, admin.NewGroupHandler, - admin.NewAccountHandler, + admin.ProvideAccountHandler, admin.NewAnnouncementHandler, admin.NewDataManagementHandler, admin.NewBackupHandler, diff --git a/backend/internal/pkg/antigravity/client.go b/backend/internal/pkg/antigravity/client.go index e318d1cdaf..39b6d2c90c 100644 --- a/backend/internal/pkg/antigravity/client.go +++ b/backend/internal/pkg/antigravity/client.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // ForbiddenError 表示上游返回 403 Forbidden @@ -279,7 +280,6 @@ func NewClient(proxyURL string) (*Client, error) { } client.Transport = transport } - return &Client{ httpClient: client, }, nil @@ -341,7 +341,7 @@ func (c *Client) ExchangeCode(ctx context.Context, code, codeVerifier string) (* } req.Header.Set("Content-Type", "application/x-www-form-urlencoded") - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("token 交换请求失败: %w", err) } @@ -383,7 +383,7 @@ func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*TokenR } req.Header.Set("Content-Type", "application/x-www-form-urlencoded") - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("token 刷新请求失败: %w", err) } @@ -414,7 +414,7 @@ func (c *Client) GetUserInfo(ctx context.Context, accessToken string) (*UserInfo } req.Header.Set("Authorization", "Bearer "+accessToken) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("用户信息请求失败: %w", err) } @@ -465,7 +465,7 @@ func (c *Client) LoadCodeAssist(ctx context.Context, accessToken string) (*LoadC req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { lastErr = fmt.Errorf("loadCodeAssist 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -544,7 +544,7 @@ func (c *Client) OnboardUser(ctx context.Context, accessToken, tierID string) (s req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { lastErr = fmt.Errorf("onboardUser 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -683,7 +683,7 @@ func (c *Client) FetchAvailableModels(ctx context.Context, accessToken, projectI req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := fetchClient.Do(req) + resp, err := servertiming.Do(fetchClient, req) if err != nil { lastErr = fmt.Errorf("fetchAvailableModels 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -842,7 +842,7 @@ func (c *Client) SetUserSettings(ctx context.Context, accessToken string) (*SetU req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1") req.Host = "daily-cloudcode-pa.googleapis.com" - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("setUserSettings 请求失败: %w", err) } @@ -885,7 +885,7 @@ func (c *Client) FetchUserInfo(ctx context.Context, accessToken, projectID strin req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1") req.Host = "daily-cloudcode-pa.googleapis.com" - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("fetchUserInfo 请求失败: %w", err) } diff --git a/backend/internal/pkg/apicompat/anthropic_responses_test.go b/backend/internal/pkg/apicompat/anthropic_responses_test.go index 8997835c2a..db6b49aa9b 100644 --- a/backend/internal/pkg/apicompat/anthropic_responses_test.go +++ b/backend/internal/pkg/apicompat/anthropic_responses_test.go @@ -718,7 +718,7 @@ func TestStreamingToolCallDoneWithoutDeltaEmitsArguments(t *testing.T) { assert.Equal(t, "content_block_stop", events[1].Type) } -func TestStreamingReadToolDropsEmptyPages(t *testing.T) { +func TestStreamingReadToolStreamsDeltas(t *testing.T) { state := NewResponsesEventToAnthropicState() ResponsesEventToAnthropicEvents(&ResponsesStreamEvent{ @@ -739,18 +739,17 @@ func TestStreamingReadToolDropsEmptyPages(t *testing.T) { OutputIndex: 0, Delta: `{"file_path":"/tmp/demo.py","limit":2000,"offset":0,"pages":""}`, }, state) - assert.Len(t, events, 0) + require.Len(t, events, 1, "Read tool deltas must be streamed like any other tool") + assert.Equal(t, "content_block_delta", events[0].Type) + assert.Equal(t, "input_json_delta", events[0].Delta.Type) events = ResponsesEventToAnthropicEvents(&ResponsesStreamEvent{ Type: "response.function_call_arguments.done", OutputIndex: 0, Arguments: `{"file_path":"/tmp/demo.py","limit":2000,"offset":0,"pages":""}`, }, state) - require.Len(t, events, 2) - assert.Equal(t, "content_block_delta", events[0].Type) - assert.Equal(t, "input_json_delta", events[0].Delta.Type) - assert.JSONEq(t, `{"file_path":"/tmp/demo.py","limit":2000,"offset":0}`, events[0].Delta.PartialJSON) - assert.Equal(t, "content_block_stop", events[1].Type) + require.Len(t, events, 1, "after streaming deltas, .done should just close the block") + assert.Equal(t, "content_block_stop", events[0].Type) } func TestStreamingReasoning(t *testing.T) { diff --git a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go index de8ab78df8..661b47cebe 100644 --- a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go +++ b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go @@ -102,9 +102,10 @@ func AnthropicToResponsesResponse(resp *AnthropicResponse) *ResponsesResponse { resp.Usage.CacheReadInputTokens + resp.Usage.CacheCreationInputTokens out.Usage = &ResponsesUsage{ - InputTokens: totalInputTokens, - OutputTokens: resp.Usage.OutputTokens, - TotalTokens: totalInputTokens + resp.Usage.OutputTokens, + InputTokens: totalInputTokens, + OutputTokens: resp.Usage.OutputTokens, + TotalTokens: totalInputTokens + resp.Usage.OutputTokens, + CacheCreationInputTokens: resp.Usage.CacheCreationInputTokens, } if resp.Usage.CacheReadInputTokens > 0 { out.Usage.InputTokensDetails = &ResponsesInputTokensDetails{ @@ -163,6 +164,8 @@ type AnthropicEventToResponsesState struct { OutputTokens int CacheReadInputTokens int CacheCreationInputTokens int + + StopReason string } // NewAnthropicEventToResponsesState returns an initialised stream state. @@ -404,7 +407,6 @@ func anthToResHandleContentBlockStop(evt *AnthropicStreamEvent, state *Anthropic } func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEventToResponsesState) []ResponsesStreamEvent { - // Update usage if evt.Usage != nil { state.OutputTokens = evt.Usage.OutputTokens if evt.Usage.InputTokens > 0 { @@ -417,6 +419,9 @@ func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEven state.CacheCreationInputTokens = evt.Usage.CacheCreationInputTokens } } + if evt.Delta != nil && evt.Delta.StopReason != "" { + state.StopReason = evt.Delta.StopReason + } return nil } @@ -427,15 +432,15 @@ func anthToResHandleMessageStop(state *AnthropicEventToResponsesState) []Respons } var events []ResponsesStreamEvent - - // Close any open item events = append(events, closeCurrentResponsesItem(state)...) - // Determine status status := "completed" var incompleteDetails *ResponsesIncompleteDetails + if state.StopReason == "max_tokens" { + status = "incomplete" + incompleteDetails = &ResponsesIncompleteDetails{Reason: "max_output_tokens"} + } - // Emit response.completed events = append(events, makeResponsesCompletedEvent(state, status, incompleteDetails)) state.CompletedSent = true return events @@ -497,9 +502,10 @@ func makeResponsesCompletedEvent( // back to match OpenAI Responses semantics where input_tokens is the total. totalInputTokens := state.InputTokens + state.CacheReadInputTokens + state.CacheCreationInputTokens usage := &ResponsesUsage{ - InputTokens: totalInputTokens, - OutputTokens: state.OutputTokens, - TotalTokens: totalInputTokens + state.OutputTokens, + InputTokens: totalInputTokens, + OutputTokens: state.OutputTokens, + TotalTokens: totalInputTokens + state.OutputTokens, + CacheCreationInputTokens: state.CacheCreationInputTokens, } if state.CacheReadInputTokens > 0 { usage.InputTokensDetails = &ResponsesInputTokensDetails{ @@ -507,15 +513,20 @@ func makeResponsesCompletedEvent( } } + eventType := "response.completed" + if status == "incomplete" { + eventType = "response.incomplete" + } + return ResponsesStreamEvent{ - Type: "response.completed", + Type: eventType, SequenceNumber: seq, Response: &ResponsesResponse{ ID: state.ResponseID, Object: "response", Model: state.Model, Status: status, - Output: []ResponsesOutput{}, // Simplified; full output tracking would add complexity + Output: []ResponsesOutput{}, Usage: usage, IncompleteDetails: incompleteDetails, }, diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go index 62b3a885cd..8aa9eab60a 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go @@ -1,6 +1,8 @@ package apicompat import ( + "crypto/sha256" + "encoding/hex" "encoding/json" "fmt" "strings" @@ -33,11 +35,30 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR if req.Reasoning != nil { out.ReasoningEffort = req.Reasoning.Effort } - if len(req.Tools) > 0 { - out.Tools = responsesToolsToChatTools(req.Tools) + effectiveTools, err := EffectiveResponsesTools(req) + if err != nil { + return nil, err } - if len(req.ToolChoice) > 0 { - out.ToolChoice = responsesToolChoiceToChatToolChoice(req.ToolChoice) + if len(effectiveTools) > 0 { + tools, err := responsesToolsToChatTools(effectiveTools) + if err != nil { + return nil, err + } + out.Tools = tools + } + // tools 全部被丢弃(如仅含 web_search/image_generation 等服务端工具)时不再转发 + // tool_choice:上游会拒绝 "'tool_choice' is only allowed when 'tools' are specified"。 + // 指向被丢弃工具的选择项同理(见 responsesToolChoiceToChatToolChoice)。 + if len(out.Tools) > 0 && len(req.ToolChoice) > 0 { + declared := make(map[string]bool, len(out.Tools)) + for _, tool := range out.Tools { + if tool.Function != nil { + declared[tool.Function.Name] = true + } + } + if tc := responsesToolChoiceToChatToolChoice(req.ToolChoice, declared); len(tc) > 0 { + out.ToolChoice = tc + } } if req.Text != nil { out.ResponseFormat = responsesTextFormatToChatResponseFormat(req.Text.Format) @@ -46,6 +67,110 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR return out, nil } +// EffectiveResponsesTools returns every client-executable tool declared by a +// Responses request. Newer Codex clients place their runtime tools in an +// input item shaped as {"type":"additional_tools","tools":[...]} instead of +// the top-level tools field. Chat-only upstreams must receive both forms. +func EffectiveResponsesTools(req *ResponsesRequest) ([]ResponsesTool, error) { + if req == nil { + return nil, nil + } + + tools := append([]ResponsesTool(nil), req.Tools...) + inputRaw := bytesTrimSpace(req.Input) + if len(inputRaw) == 0 || string(inputRaw) == "null" || inputRaw[0] != '[' { + return tools, nil + } + + var items []json.RawMessage + if err := json.Unmarshal(inputRaw, &items); err != nil { + return nil, fmt.Errorf("parse responses input for additional tools: %w", err) + } + for _, raw := range items { + raw = bytesTrimSpace(raw) + if len(raw) == 0 || raw[0] != '{' { + continue + } + var item struct { + Type string `json:"type"` + Tools []ResponsesTool `json:"tools"` + } + if err := json.Unmarshal(raw, &item); err != nil { + return nil, fmt.Errorf("parse responses additional tools item: %w", err) + } + if item.Type == "additional_tools" { + tools = append(tools, item.Tools...) + } + } + return tools, nil +} + +// CustomToolNames 收集 Responses 请求中 custom/freeform 工具的名字。chat 桥回程时 +// 需要据此把模型对这些工具的调用还原为 custom_tool_call 项(codex 只按该类型路由)。 +func CustomToolNames(tools []ResponsesTool) map[string]bool { + var out map[string]bool + for _, tool := range tools { + if tool.Type == "custom" && tool.Name != "" { + if out == nil { + out = make(map[string]bool) + } + out[tool.Name] = true + } + } + return out +} + +// NamespacedToolName 记录 namespace 子工具的原始归属(命名空间 + 裸子工具名)。 +type NamespacedToolName struct { + Namespace string + Name string +} + +// NamespaceToolNames 收集 Responses 请求中 namespace 子工具的摊平名 →(namespace, +// 子工具名)映射。chat 桥回程时需据此把模型对摊平工具的调用还原为带 namespace 字段 +// 的 function_call 项:codex 按 namespace+name 路由,平铺名会被判为 unsupported +// call;摊平名超长时带截断哈希(见 flattenNamespaceToolName),无法按字符串切分还原。 +// 摊平名撞名的请求已在转换阶段被显式拒绝(见 namespaceChildrenToChatTools), +// 此处映射不存在歧义。 +func NamespaceToolNames(tools []ResponsesTool) map[string]NamespacedToolName { + var out map[string]NamespacedToolName + for _, tool := range tools { + if tool.Type != "namespace" || tool.Name == "" { + continue + } + children := tool.Tools + if len(children) == 0 { + children = tool.Children + } + for _, child := range children { + if child.Type != "function" || child.Name == "" { + continue + } + if out == nil { + out = make(map[string]NamespacedToolName) + } + out[flattenNamespaceToolName(tool.Name, child.Name)] = NamespacedToolName{ + Namespace: tool.Name, + Name: child.Name, + } + } + } + return out +} + +// HasToolSearchTool 判断 Responses 请求是否声明了 tool_search 服务端工具。chat 桥 +// 回程时需据此把模型对代理工具的调用还原为 tool_search_call 项:codex 只在该项类型 +// 且 execution=client 时执行 tool search,同名 function_call 会因 payload 不匹配 +// 触发 fatal 中止整个 turn。 +func HasToolSearchTool(tools []ResponsesTool) bool { + for _, tool := range tools { + if tool.Type == "tool_search" { + return true + } + } + return false +} + // responsesInputToChatMessages converts a Responses request's instructions + // input[] into Chat Completions messages. It is a three-stage pipeline: // @@ -133,33 +258,68 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa if strings.TrimSpace(arguments) == "" { arguments = "{}" } + name := rawString(item["name"]) + // namespace 子工具的历史调用带 namespace 字段,需与请求方向的摊平 + // 命名(namespaceChildrenToChatTools)保持一致。 + if ns := rawString(item["namespace"]); ns != "" { + name = flattenNamespaceToolName(ns, name) + } + toolCall := ChatToolCall{ + ID: rawString(item["call_id"]), + Type: "function", + Function: ChatFunctionCall{ + Name: name, + Arguments: arguments, + }, + } + messages = appendAssistantToolCall(messages, toolCall, pendingReasoning) + pendingReasoning = "" + continue + case "tool_search_call": + // tool_search 调用的 arguments 是 JSON 对象(如 {"query": ...}), + // 原文即为降级 function 调用的 arguments 字符串。 + arguments := strings.TrimSpace(string(bytesTrimSpace(item["arguments"]))) + if s := rawString(item["arguments"]); s != "" { + arguments = s + } + if arguments == "" || arguments == "null" { + arguments = "{}" + } + toolCall := ChatToolCall{ + ID: rawString(item["call_id"]), + Type: "function", + Function: ChatFunctionCall{ + Name: toolSearchProxyName, + Arguments: arguments, + }, + } + messages = appendAssistantToolCall(messages, toolCall, pendingReasoning) + pendingReasoning = "" + continue + case "custom_tool_call": + // custom/freeform 工具的历史调用:input 自由文本包进降级 function 工具 + // 的 {"input": ...} 参数,与请求方向的工具降级(customToolInputSchema) + // 保持一致,模型才能把历史与当前工具定义对上。 + arguments, _ := json.Marshal(map[string]string{"input": rawString(item["input"])}) toolCall := ChatToolCall{ ID: rawString(item["call_id"]), Type: "function", Function: ChatFunctionCall{ Name: rawString(item["name"]), - Arguments: arguments, + Arguments: string(arguments), }, } - // Parallel tool calls arrive as consecutive function_call items and - // must share one assistant message; the matching tool replies then - // follow it. Merge into the immediately preceding assistant message. - if n := len(messages); n > 0 && messages[n-1].Role == "assistant" { - messages[n-1].ToolCalls = append(messages[n-1].ToolCalls, toolCall) - if messages[n-1].ReasoningContent == "" { - messages[n-1].ReasoningContent = pendingReasoning - } - } else { - messages = append(messages, ChatMessage{ - Role: "assistant", - ToolCalls: []ChatToolCall{toolCall}, - ReasoningContent: pendingReasoning, - }) - } + messages = appendAssistantToolCall(messages, toolCall, pendingReasoning) pendingReasoning = "" continue - case "function_call_output": - content, _ := json.Marshal(rawString(item["output"])) + case "function_call_output", "custom_tool_call_output", "tool_search_output": + outputRaw := bytesTrimSpace(item["output"]) + outputText := rawString(outputRaw) + if outputText == "" && len(outputRaw) > 0 && string(outputRaw) != "null" && string(outputRaw) != `""` { + // 对象/数组形式的输出(如 tool_search 的结果列表)整体字符串化。 + outputText = string(outputRaw) + } + content, _ := json.Marshal(outputText) messages = append(messages, ChatMessage{ Role: "tool", ToolCallID: rawString(item["call_id"]), @@ -184,9 +344,9 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa // Only genuine message items become chat messages. Codex emits other // Responses item types with no Chat equivalent (web_search_call, - // local_shell_call, custom tool calls, file_search_call, ...). Converting - // them via the generic path would insert a spurious message between an - // assistant tool_calls message and its tool reply, which DeepSeek rejects + // local_shell_call, file_search_call, ...). Converting them via the + // generic path would insert a spurious message between an assistant + // tool_calls message and its tool reply, which DeepSeek rejects // ("insufficient tool messages following tool_calls message"). Skip them. if itemType != "" && itemType != "message" { pendingReasoning = "" @@ -213,6 +373,25 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa return messages, nil } +// appendAssistantToolCall merges a tool call into the chat message list. +// Parallel tool calls arrive as consecutive *_call items and must share one +// assistant message; the matching tool replies then follow it. Merge into the +// immediately preceding assistant message. +func appendAssistantToolCall(messages []ChatMessage, toolCall ChatToolCall, pendingReasoning string) []ChatMessage { + if n := len(messages); n > 0 && messages[n-1].Role == "assistant" { + messages[n-1].ToolCalls = append(messages[n-1].ToolCalls, toolCall) + if messages[n-1].ReasoningContent == "" { + messages[n-1].ReasoningContent = pendingReasoning + } + return messages + } + return append(messages, ChatMessage{ + Role: "assistant", + ToolCalls: []ChatToolCall{toolCall}, + ReasoningContent: pendingReasoning, + }) +} + // normalizeChatMessages is the single place that enforces the tool-call // invariant the DeepSeek / OpenAI Chat Completions schema requires: an assistant // message with tool_calls must be immediately followed by one tool message per @@ -427,39 +606,187 @@ func chatContentFromSingleResponsesPart(partType string, part map[string]json.Ra } } -func responsesToolsToChatTools(tools []ResponsesTool) []ChatTool { +// customToolInputSchema 是 custom/freeform 工具降级为 function 工具时的参数 schema。 +// chat 协议无法表达 custom 工具的自由文本输入(及其 grammar 约束),退化为单一 +// input 字符串参数;回程时再从 arguments 的 input 字段还原(见 +// extractCustomToolCallInput)。 +const customToolInputSchema = `{"type":"object","properties":{"input":{"type":"string","description":"The raw input for this tool, passed through verbatim."}},"required":["input"]}` + +func responsesToolsToChatTools(tools []ResponsesTool) ([]ChatTool, error) { + // 顶层 function/custom 工具名集合:namespace 子工具摊平后与其撞名时,chat + // 上游无法按 namespace 区分调用归属。这类请求在原生 Responses 上游是合法的 + // (按 namespace+name 路由),歧义由摊平转换制造且无法消除,必须显式拒绝, + // 不能静默降级(重复声明发给上游、回程还原到错误工具)。 + topLevel := make(map[string]bool) + for _, tool := range tools { + if (tool.Type == "function" || tool.Type == "custom") && tool.Name != "" { + topLevel[tool.Name] = true + } + } + flatOwner := make(map[string]NamespacedToolName) + toolSearchDeclared := false out := make([]ChatTool, 0, len(tools)) for _, tool := range tools { - if tool.Type != "function" { + switch tool.Type { + case "function": + out = append(out, ChatTool{ + Type: "function", + Function: &ChatFunction{ + Name: tool.Name, + Description: tool.Description, + Parameters: tool.Parameters, + Strict: tool.Strict, + }, + }) + case "custom": + // codex 0.14x 的核心执行工具 exec 即为 custom 类型;丢弃它会让模型 + // 无法执行任何命令,必须降级为 function 工具透传。 + out = append(out, ChatTool{ + Type: "function", + Function: &ChatFunction{ + Name: tool.Name, + Description: tool.Description, + Parameters: json.RawMessage(customToolInputSchema), + }, + }) + case "tool_search": + // 代理不能改名(codex 的模型侧按 tool_search 这个名字调用),与客户端 + // 声明的同名工具无法区分——回程会把普通工具的调用劫持成 tool_search_call, + // 必须显式拒绝;重复声明 type=tool_search 去重即可。 + if topLevel[toolSearchProxyName] { + return nil, fmt.Errorf("built-in tool_search conflicts with a declared tool named %q; this upstream cannot disambiguate them, rename the tool", toolSearchProxyName) + } + if toolSearchDeclared { + continue + } + toolSearchDeclared = true + out = append(out, toolSearchProxyChatTool()) + case "namespace": + flattened, err := namespaceChildrenToChatTools(tool, topLevel, flatOwner) + if err != nil { + return nil, err + } + out = append(out, flattened...) + } + // 其余类型(web_search、image_generation 等服务端工具)在 chat 上游没有 + // 对应能力,维持丢弃。 + } + return out, nil +} + +// toolSearchProxyName 是 tool_search 服务端工具降级后的 function 工具名。模型对 +// 它的调用以同名 function_call 原样回传,由 codex 端路由。 +const toolSearchProxyName = "tool_search" + +const toolSearchProxySchema = `{"type":"object","properties":{"query":{"type":"string","description":"Search query for tools or connectors to load."},"limit":{"type":"integer","description":"Maximum number of tool groups to return."}},"required":["query"]}` + +func toolSearchProxyChatTool() ChatTool { + return ChatTool{ + Type: "function", + Function: &ChatFunction{ + Name: toolSearchProxyName, + Description: "Search and load Codex tools, plugins, connectors, and MCP namespaces for the current task.", + Parameters: json.RawMessage(toolSearchProxySchema), + }, + } +} + +// namespaceChildrenToChatTools 将 namespace 工具的子 function 工具摊平为顶层 +// function 工具,名字加 "__" 前缀。摊平名与顶层工具或其他 namespace +// 撞名时返回错误(歧义不可消除,显式拒绝);同一 (namespace, 子工具) 的重复声明 +// 去重后不算冲突。 +func namespaceChildrenToChatTools(tool ResponsesTool, topLevel map[string]bool, flatOwner map[string]NamespacedToolName) ([]ChatTool, error) { + if tool.Name == "" { + return nil, nil + } + children := tool.Tools + if len(children) == 0 { + children = tool.Children + } + var out []ChatTool + for _, child := range children { + if child.Type != "function" || child.Name == "" { continue } + flat := flattenNamespaceToolName(tool.Name, child.Name) + entry := NamespacedToolName{Namespace: tool.Name, Name: child.Name} + if topLevel[flat] { + return nil, fmt.Errorf("namespace tool %q/%q flattens to %q which conflicts with a top-level tool of the same name; this upstream cannot disambiguate them, rename one of the tools", tool.Name, child.Name, flat) + } + if prev, ok := flatOwner[flat]; ok { + if prev == entry { + continue + } + return nil, fmt.Errorf("namespace tools %q/%q and %q/%q both flatten to %q; this upstream cannot disambiguate them, rename one of the tools", prev.Namespace, prev.Name, tool.Name, child.Name, flat) + } + flatOwner[flat] = entry out = append(out, ChatTool{ Type: "function", Function: &ChatFunction{ - Name: tool.Name, - Description: tool.Description, - Parameters: tool.Parameters, - Strict: tool.Strict, + Name: flat, + Description: child.Description, + Parameters: child.Parameters, + Strict: child.Strict, }, }) } - return out + return out, nil } -func responsesToolChoiceToChatToolChoice(raw json.RawMessage) json.RawMessage { +// chatToolNameMaxLen 是 Chat Completions function 工具名的通用长度上限。 +const chatToolNameMaxLen = 64 + +// flattenNamespaceToolName 生成 namespace 子工具的摊平名;超长时截断并追加 +// sha256 短哈希保证唯一性。 +func flattenNamespaceToolName(namespace, name string) string { + full := namespace + "__" + name + if len(full) <= chatToolNameMaxLen { + return full + } + sum := sha256.Sum256([]byte(full)) + suffix := "__" + hex.EncodeToString(sum[:4]) + prefixLen := chatToolNameMaxLen - len(suffix) + var prefix strings.Builder + for _, ch := range full { + if prefix.Len()+len(string(ch)) > prefixLen { + break + } + _, _ = prefix.WriteRune(ch) + } + return prefix.String() + suffix +} + +// responsesToolChoiceToChatToolChoice 把 Responses 的 tool_choice 转为 chat 形态。 +// declared 是转换后实际声明的 chat 工具名集合:具名选择项仅在目标工具幸存时转发, +// 服务端工具(web_search 等)的选择项随工具本身丢弃——指向未声明工具的 tool_choice +// 会被 chat 上游 400 拒绝。返回 nil 表示丢弃 tool_choice。 +func responsesToolChoiceToChatToolChoice(raw json.RawMessage, declared map[string]bool) json.RawMessage { var choice map[string]json.RawMessage if err := json.Unmarshal(raw, &choice); err != nil { + // "auto"/"none"/"required" 等字符串形式原样转发。 return raw } - if rawString(choice["type"]) != "function" { - return raw + var name string + switch rawString(choice["type"]) { + case "tool_search": + // tool_search 未被丢弃而是降级为同名 function 代理(见 + // responsesToolsToChatTools),强制选择它同样降级为 function 选择, + // 静默丢弃会把强制搜索退化为自动选择。 + name = toolSearchProxyName + case "function", "custom": + // custom 工具已降级为 function 工具,指向它的 tool_choice 同样按 function 转换。 + name = rawString(choice["name"]) + if name == "" { + name = rawNestedString(choice["function"], "name") + } + if name == "" { + return raw + } + default: + return nil } - name := rawString(choice["name"]) - if name == "" { - name = rawNestedString(choice["function"], "name") - } - if name == "" { - return raw + if !declared[name] { + return nil } out, err := json.Marshal(map[string]any{ "type": "function", @@ -473,9 +800,38 @@ func responsesToolChoiceToChatToolChoice(raw json.RawMessage) json.RawMessage { return out } +// extractCustomToolCallInput 从降级 function 调用的 arguments 中还原 custom 工具的 +// 自由文本输入:优先取 {"input": "..."} 的 input 字段;模型未按 schema 输出时原样 +// 回传,交由客户端校验、模型重试。 +func extractCustomToolCallInput(arguments string) string { + trimmed := strings.TrimSpace(arguments) + if trimmed == "" { + return "" + } + var obj map[string]json.RawMessage + if err := json.Unmarshal([]byte(trimmed), &obj); err != nil { + return trimmed + } + if raw, ok := obj["input"]; ok { + var s string + if err := json.Unmarshal(raw, &s); err == nil { + return s + } + return trimmed + } + if len(obj) == 0 { + return "" + } + return trimmed +} + // ChatCompletionsResponseToResponses converts a non-streaming Chat Completions -// response into a Responses API response. -func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model string) *ResponsesResponse { +// response into a Responses API response. customTools 是客户端请求中 custom 工具 +// 的名字集合(见 CustomToolNames),命中的调用会还原为 custom_tool_call 项; +// toolSearch 表示客户端声明了 tool_search 工具(见 HasToolSearchTool),代理工具 +// 的调用会还原为 tool_search_call 项;namespaceTools 是 namespace 子工具的摊平名 +// 映射(见 NamespaceToolNames),命中的调用还原为带 namespace 字段的 function_call 项。 +func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model string, customTools map[string]bool, toolSearch bool, namespaceTools map[string]NamespacedToolName) *ResponsesResponse { id := "" if resp != nil { id = resp.ID @@ -500,7 +856,7 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str if len(resp.Choices) > 0 { choice := resp.Choices[0] - out.Output = chatMessageToResponsesOutput(choice.Message) + out.Output = chatMessageToResponsesOutput(choice.Message, customTools, toolSearch, namespaceTools) if choice.FinishReason == "length" { out.Status = "incomplete" out.IncompleteDetails = &ResponsesIncompleteDetails{Reason: "max_output_tokens"} @@ -515,7 +871,7 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str return out } -func chatMessageToResponsesOutput(message ChatMessage) []ResponsesOutput { +func chatMessageToResponsesOutput(message ChatMessage, customTools map[string]bool, toolSearch bool, namespaceTools map[string]NamespacedToolName) []ResponsesOutput { var outputs []ResponsesOutput if message.ReasoningContent != "" { outputs = append(outputs, ResponsesOutput{ @@ -550,6 +906,39 @@ func chatMessageToResponsesOutput(message ChatMessage) []ResponsesOutput { if strings.TrimSpace(arguments) == "" { arguments = "{}" } + if customTools[toolCall.Function.Name] { + outputs = append(outputs, ResponsesOutput{ + Type: "custom_tool_call", + ID: generateItemID(), + CallID: toolCall.ID, + Name: toolCall.Function.Name, + Input: extractCustomToolCallInput(arguments), + Status: "completed", + }) + continue + } + if toolSearch && toolCall.Function.Name == toolSearchProxyName { + outputs = append(outputs, ResponsesOutput{ + Type: "tool_search_call", + ID: generateItemID(), + CallID: toolCall.ID, + Arguments: arguments, + Status: "completed", + }) + continue + } + if ns, ok := namespaceTools[toolCall.Function.Name]; ok { + outputs = append(outputs, ResponsesOutput{ + Type: "function_call", + ID: generateItemID(), + CallID: toolCall.ID, + Name: ns.Name, + Namespace: ns.Namespace, + Arguments: arguments, + Status: "completed", + }) + continue + } outputs = append(outputs, ResponsesOutput{ Type: "function_call", ID: generateItemID(), @@ -563,6 +952,21 @@ func chatMessageToResponsesOutput(message ChatMessage) []ResponsesOutput { return outputs } +// toolSearchCallArgumentsJSON 把降级 function 调用累积的 arguments 字符串还原为 +// tool_search_call 线上要求的 JSON 对象;模型未按 schema 输出(非法 JSON)时按 +// 字符串值兜底,交由 codex 解析报错后让模型重试。 +func toolSearchCallArgumentsJSON(arguments string) json.RawMessage { + trimmed := strings.TrimSpace(arguments) + if trimmed == "" { + return json.RawMessage(`{}`) + } + if json.Valid([]byte(trimmed)) { + return json.RawMessage(trimmed) + } + fallback, _ := json.Marshal(arguments) + return fallback +} + func emptyResponsesMessageOutput() ResponsesOutput { return ResponsesOutput{ Type: "message", @@ -662,6 +1066,35 @@ type ChatCompletionsToResponsesStreamState struct { ToolItemIDs map[int]string ToolOutputIndex map[int]int + // CustomTools 是客户端请求中 custom/freeform 工具的名字集合(见 + // CustomToolNames)。命中的调用按 custom_tool_call 生命周期下发,codex 才能 + // 路由回它注册的 custom 工具。 + CustomTools map[string]bool + + // ToolSearchDeclared 表示客户端请求声明了 tool_search 工具(见 + // HasToolSearchTool)。命中的代理调用按 tool_search_call 项还原,codex 只按 + // 该项类型(且 execution=client)执行 tool search。 + ToolSearchDeclared bool + + // NamespaceTools 是 namespace 子工具的摊平名 → 原始归属映射(见 + // NamespaceToolNames)。命中的调用还原为带 namespace 字段的 function_call 项, + // codex 按 namespace+name 路由。 + NamespaceTools map[string]NamespacedToolName + + // toolIsCustom 记录每个工具调用宣告时的类型判定,保证 added/done 事件的 + // 项类型一致。 + toolIsCustom map[int]bool + + // toolIsToolSearch 记录工具调用是否判定为 tool_search 代理调用。 + toolIsToolSearch map[int]bool + + // toolNamespace 记录工具调用宣告时命中的 namespace 归属(见 NamespaceTools)。 + toolNamespace map[int]NamespacedToolName + + // toolAnnounced 记录 output_item.added 是否已发出。存在 custom 工具且名字 + // 尚未到达时延迟宣告,待名字可判定类型后再补发(见 announceChatToolItem)。 + toolAnnounced map[int]bool + FinishReason string Usage *ResponsesUsage } @@ -669,12 +1102,16 @@ type ChatCompletionsToResponsesStreamState struct { // NewChatCompletionsToResponsesStreamState returns an initialized stream state. func NewChatCompletionsToResponsesStreamState(model string) *ChatCompletionsToResponsesStreamState { return &ChatCompletionsToResponsesStreamState{ - ResponseID: generateResponsesID(), - Model: model, - Created: time.Now().Unix(), - ToolCalls: make(map[int]*ChatToolCall), - ToolItemIDs: make(map[int]string), - ToolOutputIndex: make(map[int]int), + ResponseID: generateResponsesID(), + Model: model, + Created: time.Now().Unix(), + ToolCalls: make(map[int]*ChatToolCall), + ToolItemIDs: make(map[int]string), + ToolOutputIndex: make(map[int]int), + toolIsCustom: make(map[int]bool), + toolIsToolSearch: make(map[int]bool), + toolNamespace: make(map[int]NamespacedToolName), + toolAnnounced: make(map[int]bool), } } @@ -758,19 +1195,8 @@ func ChatCompletionsChunkToResponsesEvents( copyCall.Function.Arguments = "" state.ToolCalls[idx] = ©Call stored = ©Call - itemID := generateItemID() - state.ToolItemIDs[idx] = itemID + state.ToolItemIDs[idx] = generateItemID() state.ToolOutputIndex[idx] = state.allocOutputIndex() - events = append(events, chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{ - OutputIndex: state.ToolOutputIndex[idx], - Item: &ResponsesOutput{ - Type: "function_call", - ID: itemID, - CallID: stored.ID, - Name: stored.Function.Name, - Status: "in_progress", - }, - })) } else { if toolCall.ID != "" { stored.ID = toolCall.ID @@ -779,15 +1205,22 @@ func ChatCompletionsChunkToResponsesEvents( stored.Function.Name = toolCall.Function.Name } } + events = append(events, announceChatToolItem(state, idx, stored, false)...) if toolCall.Function.Arguments != "" { stored.Function.Arguments += toolCall.Function.Arguments - events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{ - OutputIndex: state.ToolOutputIndex[idx], - ItemID: state.ToolItemIDs[idx], - Delta: toolCall.Function.Arguments, - CallID: stored.ID, - Name: stored.Function.Name, - })) + // 未宣告(名字未到)时仅累积,宣告时统一补发;custom 调用的 + // arguments 是包裹 input 的 JSON 片段,无法增量还原为自由文本 + // 输入,缓冲整份 arguments 收尾时一次性下发(见 closeChatToolItems); + // tool_search 调用同样收尾时随 output_item.done 全量下发。 + if state.toolAnnounced[idx] && !state.toolIsCustom[idx] && !state.toolIsToolSearch[idx] { + events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{ + OutputIndex: state.ToolOutputIndex[idx], + ItemID: state.ToolItemIDs[idx], + Delta: toolCall.Function.Arguments, + CallID: stored.ID, + Name: stored.Function.Name, + })) + } } } if choice.FinishReason != nil && *choice.FinishReason != "" { @@ -998,6 +1431,64 @@ func ensureChatToResponsesTextPart(state *ChatCompletionsToResponsesStreamState) })} } +// announceChatToolItem 在类型可判定时发出工具调用的 output_item.added。custom +// 工具的判定依赖名字:名字未到且请求里存在 custom 工具时延迟宣告,避免 added/done +// 的项类型不一致;force 用于流收尾,名字始终未到时按 function_call 兜底。 +func announceChatToolItem( + state *ChatCompletionsToResponsesStreamState, + idx int, + stored *ChatToolCall, + force bool, +) []ResponsesStreamEvent { + if state.toolAnnounced[idx] { + return nil + } + if !force && stored.Function.Name == "" && (len(state.CustomTools) > 0 || state.ToolSearchDeclared || len(state.NamespaceTools) > 0) { + return nil + } + state.toolAnnounced[idx] = true + isCustom := state.CustomTools[stored.Function.Name] + isToolSearch := !isCustom && state.ToolSearchDeclared && stored.Function.Name == toolSearchProxyName + state.toolIsCustom[idx] = isCustom + state.toolIsToolSearch[idx] = isToolSearch + itemType := "function_call" + if isCustom { + itemType = "custom_tool_call" + } + if isToolSearch { + itemType = "tool_search_call" + } + // namespace 子工具的调用仍按 function_call 生命周期下发,但 added/done 项要 + // 还原为裸子工具名 + namespace 字段(codex 按 namespace+name 路由)。 + itemName, itemNamespace := stored.Function.Name, "" + if ns, ok := state.NamespaceTools[stored.Function.Name]; ok && !isCustom && !isToolSearch { + state.toolNamespace[idx] = ns + itemName, itemNamespace = ns.Name, ns.Namespace + } + events := []ResponsesStreamEvent{chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{ + OutputIndex: state.ToolOutputIndex[idx], + Item: &ResponsesOutput{ + Type: itemType, + ID: state.ToolItemIDs[idx], + CallID: stored.ID, + Name: itemName, + Namespace: itemNamespace, + Status: "in_progress", + }, + })} + // 迟到宣告时补发已累积的参数增量(custom/tool_search 的输入收尾统一下发,不补发)。 + if !isCustom && !isToolSearch && stored.Function.Arguments != "" { + events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{ + OutputIndex: state.ToolOutputIndex[idx], + ItemID: state.ToolItemIDs[idx], + Delta: stored.Function.Arguments, + CallID: stored.ID, + Name: stored.Function.Name, + })) + } + return events +} + // closeChatToolItems emits function_call_arguments.done + output_item.done for // every tool call opened during the stream, carrying the full call_id/name/ // arguments so codex can deserialize and execute the call. Mirrors cc-switch's @@ -1016,17 +1507,72 @@ func closeChatToolItems(state *ChatCompletionsToResponsesStreamState) []Response if !opened { continue } + // 名字始终未到导致尚未宣告的调用,收尾前按最终名字兜底宣告。 + events = append(events, announceChatToolItem(state, i, toolCall, true)...) arguments := toolCall.Function.Arguments if strings.TrimSpace(arguments) == "" { arguments = "{}" } outputIndex := state.ToolOutputIndex[i] + if state.toolIsCustom[i] { + // custom 调用按 custom_tool_call 生命周期收尾:input 在此处一次性下发 + // (流中不产出增量,见 ChatCompletionsChunkToResponsesEvents)。 + input := extractCustomToolCallInput(arguments) + if input != "" { + events = append(events, chatToResponsesEvent(state, "response.custom_tool_call_input.delta", &ResponsesStreamEvent{ + OutputIndex: outputIndex, + ItemID: itemID, + Delta: input, + })) + } + events = append(events, + chatToResponsesEvent(state, "response.custom_tool_call_input.done", &ResponsesStreamEvent{ + OutputIndex: outputIndex, + ItemID: itemID, + CallID: toolCall.ID, + Name: toolCall.Function.Name, + Input: input, + }), + chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{ + OutputIndex: outputIndex, + Item: &ResponsesOutput{ + Type: "custom_tool_call", + ID: itemID, + CallID: toolCall.ID, + Name: toolCall.Function.Name, + Input: input, + Status: "completed", + }, + }), + ) + continue + } + if state.toolIsToolSearch[i] { + // tool_search 调用按 tool_search_call 项收尾:codex 从 output_item.done + // 物化该调用(无参数增量事件),arguments 全量随项下发。 + events = append(events, chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{ + OutputIndex: outputIndex, + Item: &ResponsesOutput{ + Type: "tool_search_call", + ID: itemID, + CallID: toolCall.ID, + Arguments: arguments, + Status: "completed", + }, + })) + continue + } + // namespace 子工具调用在宣告时已记录归属,收尾项同样带还原名与 namespace。 + name, namespace := toolCall.Function.Name, "" + if ns, ok := state.toolNamespace[i]; ok { + name, namespace = ns.Name, ns.Namespace + } events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.done", &ResponsesStreamEvent{ OutputIndex: outputIndex, ItemID: itemID, CallID: toolCall.ID, - Name: toolCall.Function.Name, + Name: name, Arguments: arguments, }), chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{ @@ -1035,7 +1581,8 @@ func closeChatToolItems(state *ChatCompletionsToResponsesStreamState) []Response Type: "function_call", ID: itemID, CallID: toolCall.ID, - Name: toolCall.Function.Name, + Name: name, + Namespace: namespace, Arguments: arguments, Status: "completed", }, @@ -1078,11 +1625,37 @@ func (state *ChatCompletionsToResponsesStreamState) chatOutput() []ResponsesOutp if strings.TrimSpace(arguments) == "" { arguments = "{}" } + if state.toolIsCustom[i] { + outputs = append(outputs, ResponsesOutput{ + Type: "custom_tool_call", + ID: generateItemID(), + CallID: toolCall.ID, + Name: toolCall.Function.Name, + Input: extractCustomToolCallInput(arguments), + Status: "completed", + }) + continue + } + if state.toolIsToolSearch[i] { + outputs = append(outputs, ResponsesOutput{ + Type: "tool_search_call", + ID: generateItemID(), + CallID: toolCall.ID, + Arguments: arguments, + Status: "completed", + }) + continue + } + name, namespace := toolCall.Function.Name, "" + if ns, ok := state.toolNamespace[i]; ok { + name, namespace = ns.Name, ns.Namespace + } outputs = append(outputs, ResponsesOutput{ Type: "function_call", ID: generateItemID(), CallID: toolCall.ID, - Name: toolCall.Function.Name, + Name: name, + Namespace: namespace, Arguments: arguments, Status: "completed", }) diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go new file mode 100644 index 0000000000..5b1d994eb3 --- /dev/null +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go @@ -0,0 +1,923 @@ +package apicompat + +// custom/freeform 工具(如 Codex 0.14x 的 exec)在 responses→chat 桥上的双向转换。 +// 背景:Codex 的核心命令执行工具 exec 是 type=custom(输入为自由文本),此前被 +// responsesToolsToChatTools 丢弃,导致模型工具列表中没有 exec、无法执行任何命令。 + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestResponsesToChatCompletionsRequest_CustomToolBecomesFunctionTool(t *testing.T) { + req := &ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"run dir"`), + Tools: []ResponsesTool{ + {Type: "custom", Name: "exec", Description: "Run JavaScript code"}, + {Type: "function", Name: "wait", Parameters: json.RawMessage(`{"type":"object","properties":{}}`)}, + }, + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.Len(t, out.Tools, 2) + + assert.Equal(t, "function", out.Tools[0].Type) + assert.Equal(t, "exec", out.Tools[0].Function.Name) + assert.Equal(t, "Run JavaScript code", out.Tools[0].Function.Description) + assert.JSONEq(t, customToolInputSchema, string(out.Tools[0].Function.Parameters)) + + assert.Equal(t, "wait", out.Tools[1].Function.Name) +} + +func TestResponsesToChatCompletionsRequest_AdditionalToolsItem(t *testing.T) { + req := &ResponsesRequest{ + Model: "gpt-test", + Input: json.RawMessage(`[ + {"type":"additional_tools","role":"developer","tools":[ + {"type":"custom","name":"exec","description":"Run PowerShell","format":{"type":"text"}}, + {"type":"function","name":"wait","parameters":{"type":"object","properties":{}}}, + {"type":"namespace","name":"collaboration","tools":[ + {"type":"function","name":"send_message","parameters":{"type":"object","properties":{}}} + ]} + ]}, + {"type":"message","role":"user","content":[{"type":"input_text","text":"run Get-Location"}]} + ]`), + ToolChoice: json.RawMessage(`"auto"`), + } + + effective, err := EffectiveResponsesTools(req) + require.NoError(t, err) + require.Len(t, effective, 3) + assert.True(t, CustomToolNames(effective)["exec"]) + assert.Equal(t, NamespacedToolName{Namespace: "collaboration", Name: "send_message"}, NamespaceToolNames(effective)["collaboration__send_message"]) + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.Len(t, out.Tools, 3) + assert.Equal(t, "exec", out.Tools[0].Function.Name) + assert.Equal(t, "wait", out.Tools[1].Function.Name) + assert.Equal(t, "collaboration__send_message", out.Tools[2].Function.Name) + assert.JSONEq(t, `"auto"`, string(out.ToolChoice)) + + require.Len(t, out.Messages, 1, "additional_tools must not become a chat message") + assert.Equal(t, "user", out.Messages[0].Role) +} + +func TestEffectiveResponsesTools_SkipsStringInputItems(t *testing.T) { + req := &ResponsesRequest{ + Input: json.RawMessage(`["plain input",{"type":"additional_tools","tools":[{"type":"custom","name":"exec"}]}]`), + } + + tools, err := EffectiveResponsesTools(req) + require.NoError(t, err) + require.Len(t, tools, 1) + assert.Equal(t, "exec", tools[0].Name) +} + +func TestResponsesToChatCompletionsRequest_DropsToolChoiceWhenNoConvertibleTools(t *testing.T) { + req := &ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "web_search"}, + {Type: "image_generation"}, + }, + ToolChoice: json.RawMessage(`"auto"`), + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + + assert.Empty(t, out.Tools) + assert.Empty(t, out.ToolChoice, "tools 为空时转发 tool_choice 会被上游 400 拒绝") +} + +func TestResponsesToChatCompletionsRequest_CustomToolChoiceMapsToFunctionChoice(t *testing.T) { + req := &ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"run dir"`), + Tools: []ResponsesTool{{Type: "custom", Name: "exec"}}, + ToolChoice: json.RawMessage(`{"type":"custom","name":"exec"}`), + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + + assert.JSONEq(t, `{"type":"function","function":{"name":"exec"}}`, string(out.ToolChoice)) +} + +func TestResponsesInputToChatMessages_CustomToolCallHistory(t *testing.T) { + input := json.RawMessage(`[ + {"role":"user","content":"list files"}, + {"type":"custom_tool_call","call_id":"call_1","name":"exec","input":"dir"}, + {"type":"custom_tool_call_output","call_id":"call_1","output":"main.go"} + ]`) + + messages, err := responsesInputToChatMessages("", input) + require.NoError(t, err) + require.Len(t, messages, 3) + + assert.Equal(t, []string{"user", "assistant", "tool"}, chatMessageRoles(messages)) + + require.Len(t, messages[1].ToolCalls, 1) + toolCall := messages[1].ToolCalls[0] + assert.Equal(t, "call_1", toolCall.ID) + assert.Equal(t, "exec", toolCall.Function.Name) + assert.JSONEq(t, `{"input":"dir"}`, toolCall.Function.Arguments) + + assert.Equal(t, "call_1", messages[2].ToolCallID) + assert.JSONEq(t, `"main.go"`, string(messages[2].Content)) +} + +func TestChatCompletionsResponseToResponses_CustomToolCallOutputItem(t *testing.T) { + resp := &ChatCompletionsResponse{ + ID: "cc-1", + Choices: []ChatChoice{{ + Message: ChatMessage{ + Role: "assistant", + ToolCalls: []ChatToolCall{ + {ID: "call_1", Function: ChatFunctionCall{Name: "exec", Arguments: `{"input": "dir"}`}}, + {ID: "call_2", Function: ChatFunctionCall{Name: "wait", Arguments: `{"cell_id": 3}`}}, + }, + }, + }}, + } + + out := ChatCompletionsResponseToResponses(resp, "glm-5.2", map[string]bool{"exec": true}, false, nil) + require.Len(t, out.Output, 2) + + assert.Equal(t, "custom_tool_call", out.Output[0].Type) + assert.Equal(t, "call_1", out.Output[0].CallID) + assert.Equal(t, "exec", out.Output[0].Name) + assert.Equal(t, "dir", out.Output[0].Input) + assert.Empty(t, out.Output[0].Arguments) + + assert.Equal(t, "function_call", out.Output[1].Type) + assert.Equal(t, "wait", out.Output[1].Name) + assert.Equal(t, `{"cell_id": 3}`, out.Output[1].Arguments) +} + +func TestExtractCustomToolCallInput_FallsBackToRawArguments(t *testing.T) { + assert.Equal(t, "dir", extractCustomToolCallInput(`{"input": "dir"}`)) + assert.Equal(t, "console.log(1)", extractCustomToolCallInput(`console.log(1)`)) + assert.Equal(t, `{"other": "x"}`, extractCustomToolCallInput(`{"other": "x"}`)) + assert.Equal(t, "", extractCustomToolCallInput(`{}`)) + assert.Equal(t, "", extractCustomToolCallInput("")) +} + +func TestChatCompletionsChunkToResponsesEvents_CustomToolCallStream(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.CustomTools = map[string]bool{"exec": true} + + idx := 0 + chunk := &ChatCompletionsChunk{ + ID: "cc-1", + Choices: []ChatChunkChoice{{ + Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{ + Index: &idx, + ID: "call_1", + Function: ChatFunctionCall{Name: "exec", Arguments: `{"input": "dir"}`}, + }}, + }, + }}, + } + + events := ChatCompletionsChunkToResponsesEvents(chunk, state) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + var added, inputDone, itemDone *ResponsesStreamEvent + for i := range events { + evt := &events[i] + switch evt.Type { + case "response.output_item.added": + if evt.Item != nil && evt.Item.Type != "message" && evt.Item.Type != "reasoning" { + added = evt + } + case "response.custom_tool_call_input.done": + inputDone = evt + case "response.output_item.done": + if evt.Item != nil && evt.Item.Type == "custom_tool_call" { + itemDone = evt + } + case "response.function_call_arguments.delta", "response.function_call_arguments.done": + t.Fatalf("custom 工具调用不应产出 function_call 参数事件: %s", evt.Type) + } + } + + require.NotNil(t, added, "缺少 custom_tool_call 的 output_item.added") + assert.Equal(t, "custom_tool_call", added.Item.Type) + assert.Equal(t, "exec", added.Item.Name) + + require.NotNil(t, inputDone, "缺少 response.custom_tool_call_input.done") + assert.Equal(t, "dir", inputDone.Input) + assert.Equal(t, "call_1", inputDone.CallID) + + require.NotNil(t, itemDone, "缺少 custom_tool_call 的 output_item.done") + assert.Equal(t, "call_1", itemDone.Item.CallID) + assert.Equal(t, "exec", itemDone.Item.Name) + assert.Equal(t, "dir", itemDone.Item.Input) + assert.Empty(t, itemDone.Item.Arguments) + + // response.completed 的 output 数组同样携带 custom_tool_call 项。 + final := events[len(events)-1] + require.Equal(t, "response.completed", final.Type) + require.NotNil(t, final.Response) + foundCustom := false + for _, item := range final.Response.Output { + if item.Type == "custom_tool_call" { + foundCustom = true + assert.Equal(t, "exec", item.Name) + assert.Equal(t, "dir", item.Input) + } + } + assert.True(t, foundCustom, "response.completed 缺少 custom_tool_call 输出项") +} + +func TestResponsesToChatCompletionsRequest_ToolSearchToolBecomesProxyFunction(t *testing.T) { + req := &ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "tool_search"}}, + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.Len(t, out.Tools, 1) + + assert.Equal(t, "function", out.Tools[0].Type) + assert.Equal(t, "tool_search", out.Tools[0].Function.Name) + assert.Contains(t, string(out.Tools[0].Function.Parameters), `"query"`) +} + +// codex 只在 ResponseItem 为 tool_search_call 变体且 execution=client 时执行 +// tool search;同名 function_call 会命中 ToolSearchHandler 后因 payload 不匹配 +// 触发 FunctionCallError::Fatal,直接中止整个 turn,因此回程必须还原项类型。 +func TestChatCompletionsResponseToResponses_ToolSearchCallOutputItem(t *testing.T) { + resp := &ChatCompletionsResponse{ + ID: "cc-1", + Choices: []ChatChoice{{ + Message: ChatMessage{ + Role: "assistant", + ToolCalls: []ChatToolCall{ + {ID: "call_s", Function: ChatFunctionCall{Name: "tool_search", Arguments: `{"query":"gmail","limit":2}`}}, + }, + }, + }}, + } + + out := ChatCompletionsResponseToResponses(resp, "glm-5.2", nil, true, nil) + require.Len(t, out.Output, 1) + + item := out.Output[0] + assert.Equal(t, "tool_search_call", item.Type) + assert.Equal(t, "call_s", item.CallID) + + // 线上形态:execution 必须为 "client"(codex 的必填字段,非 client 被忽略), + // arguments 必须是 JSON 对象而非字符串(codex 按对象解析 query/limit)。 + b, err := json.Marshal(item) + require.NoError(t, err) + var m map[string]any + require.NoError(t, json.Unmarshal(b, &m)) + assert.Equal(t, "client", m["execution"]) + args, ok := m["arguments"].(map[string]any) + require.True(t, ok, "arguments 必须序列化为 JSON 对象") + assert.Equal(t, "gmail", args["query"]) +} + +func TestChatCompletionsResponseToResponses_ToolSearchNotDeclaredKeepsFunctionCall(t *testing.T) { + resp := &ChatCompletionsResponse{ + Choices: []ChatChoice{{ + Message: ChatMessage{ + Role: "assistant", + ToolCalls: []ChatToolCall{ + {ID: "call_s", Function: ChatFunctionCall{Name: "tool_search", Arguments: `{"query":"gmail"}`}}, + }, + }, + }}, + } + + // 客户端未声明 type=tool_search 时,同名普通 function 工具不受影响。 + out := ChatCompletionsResponseToResponses(resp, "glm-5.2", nil, false, nil) + require.Len(t, out.Output, 1) + assert.Equal(t, "function_call", out.Output[0].Type) +} + +func TestChatCompletionsChunkToResponsesEvents_ToolSearchCallStream(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.ToolSearchDeclared = true + + idx := 0 + chunk := &ChatCompletionsChunk{ + ID: "cc-1", + Choices: []ChatChunkChoice{{ + Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{ + Index: &idx, + ID: "call_s", + Function: ChatFunctionCall{Name: "tool_search", Arguments: `{"query":"gmail"}`}, + }}, + }, + }}, + } + + events := ChatCompletionsChunkToResponsesEvents(chunk, state) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + var added, itemDone *ResponsesStreamEvent + for i := range events { + evt := &events[i] + switch evt.Type { + case "response.output_item.added": + if evt.Item != nil && evt.Item.Type != "message" && evt.Item.Type != "reasoning" { + added = evt + } + case "response.output_item.done": + if evt.Item != nil && evt.Item.Type == "tool_search_call" { + itemDone = evt + } + case "response.function_call_arguments.delta", "response.function_call_arguments.done", + "response.custom_tool_call_input.delta", "response.custom_tool_call_input.done": + t.Fatalf("tool_search 调用不应产出 %s", evt.Type) + } + } + + require.NotNil(t, added, "缺少 tool_search_call 的 output_item.added") + assert.Equal(t, "tool_search_call", added.Item.Type) + + require.NotNil(t, itemDone, "缺少 tool_search_call 的 output_item.done") + assert.Equal(t, "call_s", itemDone.Item.CallID) + + // SSE 线上形态经 responsesItemWire 白名单重组,必须单独断言。 + sse, err := ResponsesEventToSSE(*itemDone) + require.NoError(t, err) + assert.Contains(t, sse, `"execution":"client"`) + assert.Contains(t, sse, `"arguments":{"query":"gmail"}`) + assert.Contains(t, sse, `"call_id":"call_s"`) + + // response.completed 的 output 数组同样携带 tool_search_call 项。 + final := events[len(events)-1] + require.Equal(t, "response.completed", final.Type) + require.NotNil(t, final.Response) + found := false + for _, item := range final.Response.Output { + if item.Type == "tool_search_call" { + found = true + assert.Equal(t, "call_s", item.CallID) + } + } + assert.True(t, found, "response.completed 缺少 tool_search_call 输出项") +} + +func TestHasToolSearchTool(t *testing.T) { + assert.True(t, HasToolSearchTool([]ResponsesTool{{Type: "tool_search"}})) + assert.False(t, HasToolSearchTool([]ResponsesTool{{Type: "function", Name: "tool_search"}})) + assert.False(t, HasToolSearchTool(nil)) +} + +func TestResponsesToChatCompletionsRequest_NamespaceToolFlattensChildren(t *testing.T) { + req := &ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{ + Type: "namespace", + Name: "gmail", + Tools: []ResponsesTool{ + {Type: "function", Name: "send", Description: "Send mail", Parameters: json.RawMessage(`{"type":"object","properties":{}}`)}, + {Type: "custom", Name: "ignored_child"}, + }, + }}, + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.Len(t, out.Tools, 1, "namespace 子工具中仅 function 类型被摊平") + + assert.Equal(t, "gmail__send", out.Tools[0].Function.Name) + assert.Equal(t, "Send mail", out.Tools[0].Function.Description) +} + +func TestResponsesToolsParsing_StringToolBecomesCustom(t *testing.T) { + var req ResponsesRequest + require.NoError(t, json.Unmarshal([]byte(`{"model":"glm-5.2","input":"hi","tools":["exec",{"type":"function","name":"wait"}]}`), &req)) + + require.Len(t, req.Tools, 2) + assert.Equal(t, "custom", req.Tools[0].Type) + assert.Equal(t, "exec", req.Tools[0].Name) + assert.Equal(t, "function", req.Tools[1].Type) + + assert.True(t, CustomToolNames(req.Tools)["exec"]) +} + +func TestFlattenNamespaceToolName_CapsAt64WithHashSuffix(t *testing.T) { + assert.Equal(t, "gmail__send", flattenNamespaceToolName("gmail", "send")) + + long := flattenNamespaceToolName("very_long_namespace_prefix_for_testing_purposes", "and_a_rather_long_tool_name_too") + assert.LessOrEqual(t, len(long), 64) + assert.Contains(t, long, "__") + // 同输入结果稳定 + assert.Equal(t, long, flattenNamespaceToolName("very_long_namespace_prefix_for_testing_purposes", "and_a_rather_long_tool_name_too")) +} + +func TestResponsesInputToChatMessages_ToolSearchCallHistory(t *testing.T) { + input := json.RawMessage(`[ + {"role":"user","content":"find tools"}, + {"type":"tool_search_call","call_id":"call_s","arguments":{"query":"gmail"}}, + {"type":"tool_search_output","call_id":"call_s","output":{"groups":["gmail"]}} + ]`) + + messages, err := responsesInputToChatMessages("", input) + require.NoError(t, err) + require.Len(t, messages, 3) + + require.Len(t, messages[1].ToolCalls, 1) + assert.Equal(t, "tool_search", messages[1].ToolCalls[0].Function.Name) + assert.JSONEq(t, `{"query":"gmail"}`, messages[1].ToolCalls[0].Function.Arguments) + + assert.Equal(t, "tool", messages[2].Role) + assert.Equal(t, "call_s", messages[2].ToolCallID) + assert.JSONEq(t, `"{\"groups\":[\"gmail\"]}"`, string(messages[2].Content)) +} + +func TestResponsesInputToChatMessages_NamespacedFunctionCallHistory(t *testing.T) { + input := json.RawMessage(`[ + {"type":"function_call","call_id":"call_n","name":"send","namespace":"gmail","arguments":"{\"to\":\"a\"}"}, + {"type":"function_call_output","call_id":"call_n","output":"ok"} + ]`) + + messages, err := responsesInputToChatMessages("", input) + require.NoError(t, err) + require.Len(t, messages, 2) + + require.Len(t, messages[0].ToolCalls, 1) + assert.Equal(t, "gmail__send", messages[0].ToolCalls[0].Function.Name) +} + +func TestChatCompletionsChunkToResponsesEvents_CustomToolNameArrivesLate(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.CustomTools = map[string]bool{"exec": true} + + idx := 0 + chunk1 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, ID: "call_1", Function: ChatFunctionCall{Arguments: `{"inp`}}}, + }}}} + chunk2 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Name: "exec", Arguments: `ut": "dir"}`}}}, + }}}} + + var events []ResponsesStreamEvent + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk1, state)...) + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk2, state)...) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + addedCount := 0 + for _, evt := range events { + switch evt.Type { + case "response.output_item.added": + if evt.Item != nil && evt.Item.Type != "reasoning" && evt.Item.Type != "message" { + addedCount++ + assert.Equal(t, "custom_tool_call", evt.Item.Type, "迟到的名字命中 custom 工具时按 custom_tool_call 宣告") + assert.Equal(t, "exec", evt.Item.Name) + } + case "response.function_call_arguments.delta", "response.function_call_arguments.done": + t.Fatalf("custom 调用不应产出 function 参数事件: %s", evt.Type) + case "response.custom_tool_call_input.done": + assert.Equal(t, "dir", evt.Input) + } + } + assert.Equal(t, 1, addedCount, "工具调用只宣告一次") +} + +func TestChatCompletionsChunkToResponsesEvents_FunctionToolNameArrivesLate(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.CustomTools = map[string]bool{"exec": true} + + idx := 0 + chunk1 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, ID: "call_9", Function: ChatFunctionCall{Arguments: `{"cell`}}}, + }}}} + chunk2 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Name: "wait", Arguments: `_id": 3}`}}}, + }}}} + + var events []ResponsesStreamEvent + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk1, state)...) + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk2, state)...) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + deltas := "" + argsDone := "" + for _, evt := range events { + switch evt.Type { + case "response.function_call_arguments.delta": + deltas += evt.Delta + case "response.function_call_arguments.done": + argsDone = evt.Arguments + case "response.custom_tool_call_input.done": + t.Fatal("function 调用不应产出 custom 事件") + } + } + assert.Equal(t, `{"cell_id": 3}`, deltas, "宣告前累积的参数需在宣告时补发") + assert.Equal(t, `{"cell_id": 3}`, argsDone) +} + +// 序列化层(MarshalJSON → responsesItemWire)单独走白名单重组,事件结构体上的字段 +// 齐全不代表落到 SSE 线上的 JSON 齐全,必须在 wire 层再断言一次。 +func TestResponsesEventToSSE_CustomToolCallItemCarriesAllFields(t *testing.T) { + evt := ResponsesStreamEvent{ + Type: "response.output_item.done", + OutputIndex: 1, + Item: &ResponsesOutput{ + Type: "custom_tool_call", + ID: "item_1", + CallID: "call_1", + Name: "exec", + Input: "dir", + Status: "completed", + }, + } + + sse, err := ResponsesEventToSSE(evt) + require.NoError(t, err) + + assert.Contains(t, sse, `"call_id":"call_1"`) + assert.Contains(t, sse, `"name":"exec"`) + assert.Contains(t, sse, `"input":"dir"`) + assert.Contains(t, sse, `"type":"custom_tool_call"`) +} + +func TestNamespaceToolNames_MapsFlattenedNames(t *testing.T) { + tools := []ResponsesTool{ + {Type: "namespace", Name: "gmail", Tools: []ResponsesTool{ + {Type: "function", Name: "send"}, + {Type: "custom", Name: "skip_me"}, + }}, + {Type: "namespace", Name: "crm", Children: []ResponsesTool{ + {Type: "function", Name: "query"}, + }}, + {Type: "function", Name: "wait"}, + } + + m := NamespaceToolNames(tools) + require.Len(t, m, 2) + assert.Equal(t, NamespacedToolName{Namespace: "gmail", Name: "send"}, m["gmail__send"]) + assert.Equal(t, NamespacedToolName{Namespace: "crm", Name: "query"}, m["crm__query"]) + + // 摊平名超长时截断加哈希,无法按字符串切分还原,必须经映射反查。 + longNS := "very_long_namespace_prefix_for_testing_purposes" + longChild := "and_a_rather_long_tool_name_too" + m2 := NamespaceToolNames([]ResponsesTool{{ + Type: "namespace", Name: longNS, + Tools: []ResponsesTool{{Type: "function", Name: longChild}}, + }}) + assert.Equal(t, NamespacedToolName{Namespace: longNS, Name: longChild}, + m2[flattenNamespaceToolName(longNS, longChild)]) + + assert.Nil(t, NamespaceToolNames(nil)) +} + +// 内置 tool_search 降级后的代理 function 与客户端声明的同名工具无法区分:回程会把 +// 普通工具的调用劫持成 tool_search_call,必须显式拒绝(代理不能改名,codex 的模型 +// 侧按 tool_search 这个名字调用)。 +func TestResponsesToChatCompletionsRequest_RejectsToolSearchNameConflict(t *testing.T) { + // 与顶层 function 工具同名。 + _, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "tool_search"}, + {Type: "function", Name: "tool_search"}, + }, + }) + require.Error(t, err, "与内置 tool_search 代理撞名的 function 工具必须拒绝") + assert.Contains(t, err.Error(), "tool_search") + + // 与顶层 custom 工具同名。 + _, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "custom", Name: "tool_search"}, + {Type: "tool_search"}, + }, + }) + require.Error(t, err, "与内置 tool_search 代理撞名的 custom 工具必须拒绝") + + // 重复声明 type=tool_search 去重后只产出一个代理,不拒绝。 + out, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "tool_search"}, {Type: "tool_search"}}, + }) + require.NoError(t, err) + require.Len(t, out.Tools, 1) + assert.Equal(t, "tool_search", out.Tools[0].Function.Name) +} + +// tool_choice 指向被转换丢弃的工具(如 web_search)或不存在的名字时不能原样转发, +// chat 上游会因选择项指向未声明工具而 400;字符串形式与指向幸存工具的选择保持转发。 +func TestResponsesToChatCompletionsRequest_DropsToolChoiceForDroppedTool(t *testing.T) { + // 强制选择被丢弃的 web_search:工具没了,选择项也必须丢。 + out, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "function", Name: "wait", Parameters: json.RawMessage(`{"type":"object","properties":{}}`)}, + {Type: "web_search"}, + }, + ToolChoice: json.RawMessage(`{"type":"web_search"}`), + }) + require.NoError(t, err) + require.Len(t, out.Tools, 1) + assert.Empty(t, out.ToolChoice, "指向被丢弃服务端工具的 tool_choice 必须丢弃") + + // 具名选择指向不存在的工具名。 + out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "function", Name: "wait"}}, + ToolChoice: json.RawMessage(`{"type":"function","name":"missing"}`), + }) + require.NoError(t, err) + assert.Empty(t, out.ToolChoice, "指向不存在工具名的 tool_choice 必须丢弃") + + // 字符串形式与指向幸存工具的选择保持原有转发行为。 + out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "function", Name: "wait"}}, + ToolChoice: json.RawMessage(`"auto"`), + }) + require.NoError(t, err) + assert.JSONEq(t, `"auto"`, string(out.ToolChoice)) + + out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "function", Name: "wait"}}, + ToolChoice: json.RawMessage(`{"type":"function","name":"wait"}`), + }) + require.NoError(t, err) + assert.JSONEq(t, `{"type":"function","function":{"name":"wait"}}`, string(out.ToolChoice)) +} + +// tool_search 工具没有被丢弃而是降级为同名 function 代理,强制选择它的 tool_choice +// 必须同步降级为指向代理的 function 选择,不能静默丢弃(丢弃会把强制搜索退化为 +// 自动选择,模型可以不执行搜索)。 +func TestResponsesToChatCompletionsRequest_ToolSearchToolChoiceMapsToProxy(t *testing.T) { + out, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "tool_search"}}, + ToolChoice: json.RawMessage(`{"type":"tool_search"}`), + }) + require.NoError(t, err) + assert.JSONEq(t, `{"type":"function","function":{"name":"tool_search"}}`, string(out.ToolChoice)) + + // 未声明 type=tool_search 时强制选择它没有可指向的代理,丢弃选择项。 + out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{{Type: "function", Name: "wait"}}, + ToolChoice: json.RawMessage(`{"type":"tool_search"}`), + }) + require.NoError(t, err) + assert.Empty(t, out.ToolChoice) +} + +// 客户端请求在原生 Responses API 上合法(namespace 子工具按 namespace+name 路由), +// 是摊平转换让名字产生歧义;歧义无法消除时必须显式拒绝整个请求(400),而不是 +// 静默降级——否则重复声明发给上游、回程还原到错误工具,问题只能靠抓包定位。 +func TestResponsesToChatCompletionsRequest_RejectsAmbiguousFlattenedNames(t *testing.T) { + // 摊平名与顶层 function 工具撞名。 + _, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "function", Name: "gmail__send"}, + {Type: "namespace", Name: "gmail", Tools: []ResponsesTool{{Type: "function", Name: "send"}}}, + }, + }) + require.Error(t, err, "与顶层工具撞名的摊平必须拒绝") + assert.Contains(t, err.Error(), "gmail__send") + + // 不同 namespace 组合产生相同摊平名。 + _, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "namespace", Name: "a", Tools: []ResponsesTool{{Type: "function", Name: "b__c"}}}, + {Type: "namespace", Name: "a__b", Tools: []ResponsesTool{{Type: "function", Name: "c"}}}, + }, + }) + require.Error(t, err, "跨 namespace 撞名的摊平必须拒绝") + assert.Contains(t, err.Error(), "a__b__c") +} + +// 完全相同的 (namespace, 子工具) 重复声明不构成歧义:去重后正常转换,不拒绝。 +func TestResponsesToChatCompletionsRequest_DedupesIdenticalNamespaceChildren(t *testing.T) { + out, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "namespace", Name: "gmail", Tools: []ResponsesTool{ + {Type: "function", Name: "send"}, + {Type: "function", Name: "send"}, + }}, + }, + }) + require.NoError(t, err) + require.Len(t, out.Tools, 1, "重复声明的同一子工具只声明一次") + assert.Equal(t, "gmail__send", out.Tools[0].Function.Name) +} + +// codex 按 namespace+name 路由 namespace 子工具的调用:回程必须把摊平名还原为 +// 裸子工具名并带独立 namespace 字段,平铺名的 function_call 会被 codex 判为 +// unsupported call 拒绝执行。 +func TestChatCompletionsResponseToResponses_NamespacedToolCallRestored(t *testing.T) { + resp := &ChatCompletionsResponse{ + ID: "cc-1", + Choices: []ChatChoice{{ + Message: ChatMessage{ + Role: "assistant", + ToolCalls: []ChatToolCall{ + {ID: "call_n", Function: ChatFunctionCall{Name: "mcp__svc__echo", Arguments: `{"text":"hi"}`}}, + {ID: "call_9", Function: ChatFunctionCall{Name: "wait", Arguments: `{"cell_id": 3}`}}, + }, + }, + }}, + } + nsTools := map[string]NamespacedToolName{ + "mcp__svc__echo": {Namespace: "mcp__svc", Name: "echo"}, + } + + out := ChatCompletionsResponseToResponses(resp, "glm-5.2", nil, false, nsTools) + require.Len(t, out.Output, 2) + + item := out.Output[0] + assert.Equal(t, "function_call", item.Type) + assert.Equal(t, "echo", item.Name) + assert.Equal(t, "mcp__svc", item.Namespace) + assert.Equal(t, "call_n", item.CallID) + assert.Equal(t, `{"text":"hi"}`, item.Arguments) + + // 非流式响应体走 ResponsesOutput.MarshalJSON,namespace 必须落到线上 JSON。 + b, err := json.Marshal(item) + require.NoError(t, err) + assert.Contains(t, string(b), `"namespace":"mcp__svc"`) + assert.Contains(t, string(b), `"name":"echo"`) + + // 未命中映射的普通 function 调用不受影响,且不携带 namespace 字段。 + assert.Equal(t, "wait", out.Output[1].Name) + assert.Empty(t, out.Output[1].Namespace) + b2, err := json.Marshal(out.Output[1]) + require.NoError(t, err) + assert.NotContains(t, string(b2), `"namespace"`) +} + +func TestChatCompletionsChunkToResponsesEvents_NamespacedToolCallStream(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.NamespaceTools = map[string]NamespacedToolName{ + "mcp__svc__echo": {Namespace: "mcp__svc", Name: "echo"}, + } + + idx := 0 + chunk := &ChatCompletionsChunk{ + ID: "cc-1", + Choices: []ChatChunkChoice{{ + Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{ + Index: &idx, + ID: "call_n", + Function: ChatFunctionCall{Name: "mcp__svc__echo", Arguments: `{"text":"hi"}`}, + }}, + }, + }}, + } + + events := ChatCompletionsChunkToResponsesEvents(chunk, state) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + var added, itemDone *ResponsesStreamEvent + for i := range events { + evt := &events[i] + switch evt.Type { + case "response.output_item.added": + if evt.Item != nil && evt.Item.Type != "message" && evt.Item.Type != "reasoning" { + added = evt + } + case "response.output_item.done": + if evt.Item != nil && evt.Item.Type == "function_call" { + itemDone = evt + } + case "response.custom_tool_call_input.delta", "response.custom_tool_call_input.done": + t.Fatalf("namespace 子工具调用不应产出 custom 事件: %s", evt.Type) + } + } + + require.NotNil(t, added, "缺少 namespace 调用的 output_item.added") + assert.Equal(t, "function_call", added.Item.Type) + assert.Equal(t, "echo", added.Item.Name) + assert.Equal(t, "mcp__svc", added.Item.Namespace) + + require.NotNil(t, itemDone, "缺少 namespace 调用的 output_item.done") + assert.Equal(t, "call_n", itemDone.Item.CallID) + assert.Equal(t, "echo", itemDone.Item.Name) + assert.Equal(t, "mcp__svc", itemDone.Item.Namespace) + assert.Equal(t, `{"text":"hi"}`, itemDone.Item.Arguments) + + // SSE 线上形态经 responsesItemWire 白名单重组,必须单独断言 namespace 落线。 + sse, err := ResponsesEventToSSE(*itemDone) + require.NoError(t, err) + assert.Contains(t, sse, `"namespace":"mcp__svc"`) + assert.Contains(t, sse, `"name":"echo"`) + assert.Contains(t, sse, `"call_id":"call_n"`) + + // response.completed 的 output 数组同样携带还原后的 namespace 调用项。 + final := events[len(events)-1] + require.Equal(t, "response.completed", final.Type) + require.NotNil(t, final.Response) + found := false + for _, item := range final.Response.Output { + if item.Type == "function_call" { + found = true + assert.Equal(t, "echo", item.Name) + assert.Equal(t, "mcp__svc", item.Namespace) + } + } + assert.True(t, found, "response.completed 缺少还原后的 namespace 调用项") +} + +func TestChatCompletionsChunkToResponsesEvents_NamespacedToolNameArrivesLate(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.NamespaceTools = map[string]NamespacedToolName{ + "mcp__svc__echo": {Namespace: "mcp__svc", Name: "echo"}, + } + + idx := 0 + chunk1 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, ID: "call_n", Function: ChatFunctionCall{Arguments: `{"te`}}}, + }}}} + chunk2 := &ChatCompletionsChunk{Choices: []ChatChunkChoice{{Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Name: "mcp__svc__echo", Arguments: `xt":"hi"}`}}}, + }}}} + + var events []ResponsesStreamEvent + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk1, state)...) + events = append(events, ChatCompletionsChunkToResponsesEvents(chunk2, state)...) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + addedCount := 0 + deltas := "" + for _, evt := range events { + switch evt.Type { + case "response.output_item.added": + if evt.Item != nil && evt.Item.Type != "reasoning" && evt.Item.Type != "message" { + addedCount++ + assert.Equal(t, "echo", evt.Item.Name, "迟到的名字命中 namespace 映射时按还原名宣告") + assert.Equal(t, "mcp__svc", evt.Item.Namespace) + } + case "response.function_call_arguments.delta": + deltas += evt.Delta + } + } + assert.Equal(t, 1, addedCount, "工具调用只宣告一次") + assert.Equal(t, `{"text":"hi"}`, deltas, "宣告前累积的参数需在宣告时补发") +} + +func TestChatCompletionsChunkToResponsesEvents_FunctionToolStreamUnaffected(t *testing.T) { + state := NewChatCompletionsToResponsesStreamState("glm-5.2") + state.CustomTools = map[string]bool{"exec": true} + + idx := 0 + chunk := &ChatCompletionsChunk{ + Choices: []ChatChunkChoice{{ + Delta: ChatDelta{ + ToolCalls: []ChatToolCall{{ + Index: &idx, + ID: "call_9", + Function: ChatFunctionCall{Name: "wait", Arguments: `{"cell_id": 3}`}, + }}, + }, + }}, + } + + events := ChatCompletionsChunkToResponsesEvents(chunk, state) + events = append(events, FinalizeChatCompletionsResponsesStream(state)...) + + sawArgsDelta := false + for _, evt := range events { + if evt.Type == "response.function_call_arguments.delta" { + sawArgsDelta = true + } + if evt.Type == "response.custom_tool_call_input.done" { + t.Fatal("function 工具不应产出 custom_tool_call 事件") + } + } + assert.True(t, sawArgsDelta, "function 工具应保持原有参数增量事件") +} diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index 4075f57791..4e319f9751 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -459,7 +459,7 @@ func TestChatCompletionsResponseToResponses_DeepSeekReasoningOnlyFallsBackToMess }}, } - out := ChatCompletionsResponseToResponses(resp, "deepseek-reasoner") + out := ChatCompletionsResponseToResponses(resp, "deepseek-reasoner", nil, false, nil) require.Len(t, out.Output, 2) require.Equal(t, "reasoning", out.Output[0].Type) @@ -493,7 +493,7 @@ func TestChatCompletionsResponseToResponses_DeepSeekReasoningToolCallDoesNotFall }}, } - out := ChatCompletionsResponseToResponses(resp, "deepseek-reasoner") + out := ChatCompletionsResponseToResponses(resp, "deepseek-reasoner", nil, false, nil) require.Len(t, out.Output, 2) require.Equal(t, "reasoning", out.Output[0].Type) diff --git a/backend/internal/pkg/apicompat/responses_anthropic_cache_creation_test.go b/backend/internal/pkg/apicompat/responses_anthropic_cache_creation_test.go new file mode 100644 index 0000000000..b8c7916d2d --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_anthropic_cache_creation_test.go @@ -0,0 +1,99 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAnthropicUsageFromResponsesUsage_CacheCreation(t *testing.T) { + usage := &ResponsesUsage{ + InputTokens: 20, + OutputTokens: 5, + CacheCreationInputTokens: 6, + InputTokensDetails: &ResponsesInputTokensDetails{ + CachedTokens: 4, + }, + } + + got := anthropicUsageFromResponsesUsage(usage) + + assert.Equal(t, 10, got.InputTokens, "input = total(20) - cache_read(4) - cache_creation(6)") + assert.Equal(t, 5, got.OutputTokens) + assert.Equal(t, 4, got.CacheReadInputTokens) + assert.Equal(t, 6, got.CacheCreationInputTokens, "cache creation must be preserved") +} + +func TestAnthropicUsageFromResponsesUsage_NoCacheCreation(t *testing.T) { + usage := &ResponsesUsage{ + InputTokens: 10, + OutputTokens: 5, + InputTokensDetails: &ResponsesInputTokensDetails{ + CachedTokens: 3, + }, + } + + got := anthropicUsageFromResponsesUsage(usage) + + assert.Equal(t, 7, got.InputTokens) + assert.Equal(t, 3, got.CacheReadInputTokens) + assert.Equal(t, 0, got.CacheCreationInputTokens) +} + +func TestResponsesEventToAnthropicEvents_StreamingCacheCreation(t *testing.T) { + state := NewResponsesEventToAnthropicState() + state.MessageStartSent = true + + completedEvt := &ResponsesStreamEvent{ + Type: "response.completed", + Response: &ResponsesResponse{ + Status: "completed", + Usage: &ResponsesUsage{ + InputTokens: 20, + OutputTokens: 5, + CacheCreationInputTokens: 6, + InputTokensDetails: &ResponsesInputTokensDetails{ + CachedTokens: 4, + }, + }, + }, + } + + events := ResponsesEventToAnthropicEvents(completedEvt, state) + + var deltaEvt *AnthropicStreamEvent + for i := range events { + if events[i].Type == "message_delta" { + deltaEvt = &events[i] + break + } + } + require.NotNil(t, deltaEvt, "should have message_delta event") + require.NotNil(t, deltaEvt.Usage) + assert.Equal(t, 6, deltaEvt.Usage.CacheCreationInputTokens, "streaming cache_creation must be preserved") + assert.Equal(t, 10, deltaEvt.Usage.InputTokens, "input = 20 - 4(read) - 6(creation)") + assert.Equal(t, 4, deltaEvt.Usage.CacheReadInputTokens) +} + +func TestAnthropicToResponsesResponse_CacheCreation(t *testing.T) { + resp := AnthropicResponse{ + ID: "msg_test", + Type: "message", + Role: "assistant", + Model: "claude-opus-4-6", + Usage: AnthropicUsage{ + InputTokens: 10, + OutputTokens: 5, + CacheReadInputTokens: 4, + CacheCreationInputTokens: 6, + }, + StopReason: "end_turn", + } + + out := AnthropicToResponsesResponse(&resp) + + require.NotNil(t, out.Usage) + assert.Equal(t, 20, out.Usage.InputTokens, "total = input(10) + cache_read(4) + cache_creation(6)") + assert.Equal(t, 6, out.Usage.CacheCreationInputTokens, "cache creation must round-trip") +} diff --git a/backend/internal/pkg/apicompat/responses_namespace.go b/backend/internal/pkg/apicompat/responses_namespace.go new file mode 100644 index 0000000000..a5549760c5 --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_namespace.go @@ -0,0 +1,212 @@ +package apicompat + +import ( + "bytes" + "encoding/json" + "fmt" + "strings" +) + +// ResponsesNamespaceName identifies a function child in a Responses namespace. +// It aliases the chat bridge mapping so both native and bridged paths share one +// namespace identity contract. +type ResponsesNamespaceName = NamespacedToolName + +// FlattenResponsesNamespaces converts Codex private namespace declarations into +// public Responses function tools and rewrites namespace-qualified request calls. +func FlattenResponsesNamespaces(req map[string]any) (map[string]ResponsesNamespaceName, bool, error) { + return FlattenResponsesNamespacesExcept(req, nil) +} + +// FlattenResponsesNamespacesExcept is FlattenResponsesNamespaces with a set of +// service-owned namespace names that must remain native in the request. +func FlattenResponsesNamespacesExcept(req map[string]any, preserved map[string]bool) (map[string]ResponsesNamespaceName, bool, error) { + if req == nil { + return nil, false, nil + } + tools, ok := req["tools"].([]any) + if !ok || len(tools) == 0 { + return nil, false, nil + } + + topLevel := make(map[string]bool) + for _, raw := range tools { + tool, ok := raw.(map[string]any) + if !ok { + continue + } + typ := strings.TrimSpace(stringValue(tool["type"])) + name := strings.TrimSpace(stringValue(tool["name"])) + if (typ == "function" || typ == "custom") && name != "" { + topLevel[name] = true + } + } + + names := make(map[string]ResponsesNamespaceName) + for _, raw := range tools { + tool, ok := raw.(map[string]any) + if !ok || strings.TrimSpace(stringValue(tool["type"])) != "namespace" { + continue + } + namespace := strings.TrimSpace(stringValue(tool["name"])) + if namespace == "" || preserved[namespace] { + continue + } + for _, rawChild := range namespaceChildren(tool) { + child, ok := rawChild.(map[string]any) + if !ok || strings.TrimSpace(stringValue(child["type"])) != "function" { + continue + } + name := strings.TrimSpace(stringValue(child["name"])) + if name == "" { + continue + } + flat := flattenNamespaceToolName(namespace, name) + entry := ResponsesNamespaceName{Namespace: namespace, Name: name} + if topLevel[flat] { + return nil, false, fmt.Errorf("namespace tool %q/%q flattens to %q which conflicts with a top-level tool of the same name; this upstream cannot disambiguate them, rename one of the tools", namespace, name, flat) + } + if prev, exists := names[flat]; exists && prev != entry { + return nil, false, fmt.Errorf("namespace tools %q/%q and %q/%q both flatten to %q; this upstream cannot disambiguate them, rename one of the tools", prev.Namespace, prev.Name, namespace, name, flat) + } + names[flat] = entry + } + } + if len(names) == 0 { + return nil, false, nil + } + + flattened := make([]any, 0, len(tools)+len(names)) + seen := make(map[string]bool) + for _, raw := range tools { + tool, ok := raw.(map[string]any) + if !ok || strings.TrimSpace(stringValue(tool["type"])) != "namespace" { + flattened = append(flattened, raw) + continue + } + namespace := strings.TrimSpace(stringValue(tool["name"])) + if preserved[namespace] { + flattened = append(flattened, raw) + continue + } + for _, rawChild := range namespaceChildren(tool) { + child, ok := rawChild.(map[string]any) + if !ok || strings.TrimSpace(stringValue(child["type"])) != "function" { + continue + } + name := strings.TrimSpace(stringValue(child["name"])) + flat := flattenNamespaceToolName(namespace, name) + if name == "" || seen[flat] { + continue + } + seen[flat] = true + flatChild := make(map[string]any, len(child)) + for key, value := range child { + flatChild[key] = value + } + flatChild["name"] = flat + flattened = append(flattened, flatChild) + } + } + req["tools"] = flattened + rewriteNamespaceQualifiedCalls(req["input"], names) + if choice, ok := req["tool_choice"].(map[string]any); ok { + choiceNamespace := strings.TrimSpace(stringValue(choice["name"])) + if strings.TrimSpace(stringValue(choice["type"])) == "namespace" && !preserved[choiceNamespace] { + req["tool_choice"] = "auto" + } else { + rewriteNamespaceQualifiedCall(choice, names) + } + } + return names, true, nil +} + +// RestoreResponsesNamespaceCalls restores flattened function calls in a JSON +// Responses payload to the namespace/name identity expected by Codex. +func RestoreResponsesNamespaceCalls(payload []byte, names map[string]ResponsesNamespaceName) ([]byte, bool, error) { + if len(payload) == 0 || len(names) == 0 { + return payload, false, nil + } + var value any + if err := json.Unmarshal(payload, &value); err != nil { + return payload, false, err + } + changed := restoreResponsesNamespaceValue(value, names) + if !changed { + return payload, false, nil + } + var rebuilt bytes.Buffer + encoder := json.NewEncoder(&rebuilt) + encoder.SetEscapeHTML(false) + if err := encoder.Encode(value); err != nil { + return payload, false, err + } + return bytes.TrimSuffix(rebuilt.Bytes(), []byte("\n")), true, nil +} + +func namespaceChildren(tool map[string]any) []any { + if children, ok := tool["tools"].([]any); ok && len(children) > 0 { + return children + } + children, _ := tool["children"].([]any) + return children +} + +func rewriteNamespaceQualifiedCalls(value any, names map[string]ResponsesNamespaceName) { + switch typed := value.(type) { + case []any: + for _, item := range typed { + rewriteNamespaceQualifiedCalls(item, names) + } + case map[string]any: + if strings.TrimSpace(stringValue(typed["type"])) == "function_call" { + rewriteNamespaceQualifiedCall(typed, names) + } + for _, child := range typed { + rewriteNamespaceQualifiedCalls(child, names) + } + } +} + +func rewriteNamespaceQualifiedCall(item map[string]any, names map[string]ResponsesNamespaceName) bool { + namespace := strings.TrimSpace(stringValue(item["namespace"])) + name := strings.TrimSpace(stringValue(item["name"])) + if namespace == "" || name == "" { + return false + } + flat := flattenNamespaceToolName(namespace, name) + entry, ok := names[flat] + if !ok || entry.Namespace != namespace || entry.Name != name { + return false + } + item["name"] = flat + delete(item, "namespace") + return true +} + +func restoreResponsesNamespaceValue(value any, names map[string]ResponsesNamespaceName) bool { + changed := false + switch typed := value.(type) { + case []any: + for _, item := range typed { + changed = restoreResponsesNamespaceValue(item, names) || changed + } + case map[string]any: + if strings.TrimSpace(stringValue(typed["type"])) == "function_call" { + if entry, ok := names[strings.TrimSpace(stringValue(typed["name"]))]; ok { + typed["name"] = entry.Name + typed["namespace"] = entry.Namespace + changed = true + } + } + for _, child := range typed { + changed = restoreResponsesNamespaceValue(child, names) || changed + } + } + return changed +} + +func stringValue(value any) string { + text, _ := value.(string) + return text +} diff --git a/backend/internal/pkg/apicompat/responses_namespace_test.go b/backend/internal/pkg/apicompat/responses_namespace_test.go new file mode 100644 index 0000000000..4c8e313c5d --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_namespace_test.go @@ -0,0 +1,164 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestFlattenResponsesNamespaces_RewritesDeclarationHistoryAndChoice(t *testing.T) { + req := map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{"type": "function", "name": "plain", "description": "keep"}, + map[string]any{ + "type": "namespace", + "name": "collaboration", + "tools": []any{ + map[string]any{"type": "function", "name": "spawn_agent", "description": "spawn", "parameters": map[string]any{"type": "object"}}, + }, + }, + }, + "tool_choice": map[string]any{"type": "function", "name": "spawn_agent", "namespace": "collaboration"}, + "input": []any{ + map[string]any{"type": "function_call", "call_id": "call_1", "name": "spawn_agent", "namespace": "collaboration", "arguments": "{}"}, + map[string]any{"type": "message", "role": "user", "content": "hi", "name": "spawn_agent", "namespace": "collaboration"}, + }, + } + + names, changed, err := FlattenResponsesNamespaces(req) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, ResponsesNamespaceName{Namespace: "collaboration", Name: "spawn_agent"}, names["collaboration__spawn_agent"]) + + tools, ok := req["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 2) + plainTool, ok := tools[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "plain", plainTool["name"]) + flatTool, ok := tools[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "collaboration__spawn_agent", flatTool["name"]) + require.Equal(t, "spawn", flatTool["description"]) + + choice, ok := req["tool_choice"].(map[string]any) + require.True(t, ok) + require.Equal(t, "collaboration__spawn_agent", choice["name"]) + require.NotContains(t, choice, "namespace") + + input, ok := req["input"].([]any) + require.True(t, ok) + require.Len(t, input, 2) + call, ok := input[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "collaboration__spawn_agent", call["name"]) + require.NotContains(t, call, "namespace") + message, ok := input[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "spawn_agent", message["name"]) + require.Equal(t, "collaboration", message["namespace"]) + require.Equal(t, "gpt-5.5", req["model"]) +} + +func TestFlattenResponsesNamespaces_RejectsFlatNameCollision(t *testing.T) { + req := map[string]any{"tools": []any{ + map[string]any{"type": "function", "name": "collaboration__spawn_agent"}, + map[string]any{"type": "namespace", "name": "collaboration", "tools": []any{ + map[string]any{"type": "function", "name": "spawn_agent"}, + }}, + }} + + _, _, err := FlattenResponsesNamespaces(req) + require.ErrorContains(t, err, "conflicts with a top-level tool") +} + +func TestFlattenResponsesNamespaces_NamespaceGroupChoiceFallsBackToAuto(t *testing.T) { + req := map[string]any{ + "tools": []any{map[string]any{ + "type": "namespace", "name": "collaboration", "tools": []any{ + map[string]any{"type": "function", "name": "spawn_agent"}, + map[string]any{"type": "function", "name": "send_message"}, + }, + }}, + "tool_choice": map[string]any{"type": "namespace", "name": "collaboration"}, + } + + _, changed, err := FlattenResponsesNamespaces(req) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "auto", req["tool_choice"]) +} + +func TestFlattenResponsesNamespacesExcept_PreservesBuiltInNamespaceAndChoice(t *testing.T) { + req := map[string]any{ + "tools": []any{ + map[string]any{"type": "namespace", "name": "image_gen", "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }}, + map[string]any{"type": "namespace", "name": "collaboration", "tools": []any{ + map[string]any{"type": "function", "name": "spawn_agent"}, + }}, + }, + "tool_choice": map[string]any{"type": "namespace", "name": "image_gen"}, + } + + names, changed, err := FlattenResponsesNamespacesExcept(req, map[string]bool{"image_gen": true}) + require.NoError(t, err) + require.True(t, changed) + require.Contains(t, names, "collaboration__spawn_agent") + tools, ok := req["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 2) + preservedTool, ok := tools[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "namespace", preservedTool["type"]) + require.Equal(t, "image_gen", preservedTool["name"]) + flatTool, ok := tools[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "function", flatTool["type"]) + require.Equal(t, "collaboration__spawn_agent", flatTool["name"]) + require.Equal(t, map[string]any{"type": "namespace", "name": "image_gen"}, req["tool_choice"]) +} + +func TestFlattenResponsesNamespaces_RejectsNamespaceCollision(t *testing.T) { + req := map[string]any{"tools": []any{ + map[string]any{"type": "namespace", "name": "a", "tools": []any{ + map[string]any{"type": "function", "name": "b__c"}, + }}, + map[string]any{"type": "namespace", "name": "a__b", "tools": []any{ + map[string]any{"type": "function", "name": "c"}, + }}, + }} + + _, _, err := FlattenResponsesNamespaces(req) + require.ErrorContains(t, err, "both flatten") +} + +func TestRestoreResponsesNamespaceCalls_RewritesOnlyFunctionCalls(t *testing.T) { + payload := []byte(`{"type":"response.completed","response":{"output":[{"type":"function_call","name":"collaboration__spawn_agent","call_id":"call_1","arguments":"{}","extra":"keep"},{"type":"function_call","name":"plain","arguments":"{}"},{"type":"message","name":"collaboration__spawn_agent","content":"&value"}]}}`) + names := map[string]ResponsesNamespaceName{ + "collaboration__spawn_agent": {Namespace: "collaboration", Name: "spawn_agent"}, + } + + got, changed, err := RestoreResponsesNamespaceCalls(payload, names) + require.NoError(t, err) + require.True(t, changed) + require.JSONEq(t, `{"type":"response.completed","response":{"output":[{"type":"function_call","name":"spawn_agent","namespace":"collaboration","call_id":"call_1","arguments":"{}","extra":"keep"},{"type":"function_call","name":"plain","arguments":"{}"},{"type":"message","name":"collaboration__spawn_agent","content":"&value"}]}}`, string(got)) + require.Contains(t, string(got), "&value") + require.NotContains(t, string(got), `\u003c`) +} + +func TestRestoreResponsesNamespaceCalls_RewritesLifecycleItems(t *testing.T) { + for _, eventType := range []string{"response.output_item.added", "response.output_item.done"} { + t.Run(eventType, func(t *testing.T) { + payload := []byte(`{"type":"` + eventType + `","item":{"type":"function_call","name":"collaboration__spawn_agent","arguments":"{}"}}`) + got, changed, err := RestoreResponsesNamespaceCalls(payload, map[string]ResponsesNamespaceName{ + "collaboration__spawn_agent": {Namespace: "collaboration", Name: "spawn_agent"}, + }) + require.NoError(t, err) + require.True(t, changed) + require.JSONEq(t, `{"type":"`+eventType+`","item":{"type":"function_call","name":"spawn_agent","namespace":"collaboration","arguments":"{}"}}`, string(got)) + }) + } +} diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire.go b/backend/internal/pkg/apicompat/responses_stream_event_wire.go index df7a82e393..c2ebcd3eee 100644 --- a/backend/internal/pkg/apicompat/responses_stream_event_wire.go +++ b/backend/internal/pkg/apicompat/responses_stream_event_wire.go @@ -86,6 +86,23 @@ func (e ResponsesStreamEvent) MarshalJSON() ([]byte, error) { } return json.Marshal(m) + case "response.custom_tool_call_input.delta", "response.custom_tool_call_input.done": + m := e.wireBase() + e.putItemID(m) + m["output_index"] = e.OutputIndex + if e.CallID != "" { + m["call_id"] = e.CallID + } + if e.Name != "" { + m["name"] = e.Name + } + if e.Type == "response.custom_tool_call_input.done" { + m["input"] = e.Input + } else { + m["delta"] = e.Delta + } + return json.Marshal(m) + default: // response.created / completed / done / failed / incomplete and any // event type not shaped above keep the default struct marshalling. @@ -167,6 +184,23 @@ func responsesItemWire(item *ResponsesOutput) map[string]any { m["call_id"] = item.CallID m["name"] = item.Name m["arguments"] = item.Arguments + // namespace 子工具的还原调用:codex 按 namespace+name 路由,缺少该字段 + // 会被判为 unsupported call。 + if item.Namespace != "" { + m["namespace"] = item.Namespace + } + case "custom_tool_call": + // custom/freeform 工具调用(如 codex 的 exec):input 为自由文本。缺少 + // call_id/name 时 codex 无法路由该调用(表现为 unsupported call)。 + m["call_id"] = item.CallID + m["name"] = item.Name + m["input"] = item.Input + case "tool_search_call": + // tool_search 调用还原项:execution 必须为 "client"(否则 codex 忽略该 + // 调用),arguments 在线上是 JSON 对象而非字符串。 + m["call_id"] = item.CallID + m["execution"] = "client" + m["arguments"] = toolSearchCallArgumentsJSON(item.Arguments) } return m } diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go index b4f6871d5f..fb138a1469 100644 --- a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go +++ b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go @@ -102,6 +102,26 @@ func TestWire_ArgumentsDonePresentEvenEmpty(t *testing.T) { require.Equal(t, "", m["arguments"]) } +// TestWire_CustomToolCallInputIndexPresentAtZero guards the omitempty trap for +// custom_tool_call_input.delta/done: output_index must serialize even when 0 +// (custom tool call as the first output item). +func TestWire_CustomToolCallInputIndexPresentAtZero(t *testing.T) { + d := marshalEvent(t, ResponsesStreamEvent{ + Type: "response.custom_tool_call_input.delta", OutputIndex: 0, ItemID: "ct_1", Delta: "dir", + }) + require.Contains(t, d, "output_index") + require.EqualValues(t, 0, d["output_index"]) + require.Equal(t, "dir", d["delta"]) + + done := marshalEvent(t, ResponsesStreamEvent{ + Type: "response.custom_tool_call_input.done", OutputIndex: 0, ItemID: "ct_1", CallID: "call_1", Name: "exec", Input: "dir", + }) + require.Contains(t, done, "output_index") + require.EqualValues(t, 0, done["output_index"]) + require.Equal(t, "dir", done["input"]) + require.NotContains(t, done, "delta") +} + // TestWire_UnknownEventFallsBackToDefault ensures non-streamed event types keep // default marshalling (the response object is preserved). func TestWire_UnknownEventFallsBackToDefault(t *testing.T) { @@ -111,3 +131,56 @@ func TestWire_UnknownEventFallsBackToDefault(t *testing.T) { }) require.Contains(t, m, "response") } + +func TestResponsesOutputUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var item ResponsesOutput + require.NoError(t, json.Unmarshal([]byte(`{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "execution":"client", + "arguments":{"query":"gmail","limit":2} + }`), &item)) + require.Equal(t, "tool_search_call", item.Type) + require.Equal(t, `{"query":"gmail","limit":2}`, item.Arguments) + + wire, err := json.Marshal(item) + require.NoError(t, err) + var decoded map[string]any + require.NoError(t, json.Unmarshal(wire, &decoded)) + args, ok := decoded["arguments"].(map[string]any) + require.True(t, ok, "tool_search_call arguments must remain an object") + require.Equal(t, "gmail", args["query"]) +} + +func TestResponsesResponseUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var response ResponsesResponse + require.NoError(t, json.Unmarshal([]byte(`{ + "id":"response_1", + "object":"response", + "status":"completed", + "output":[{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "arguments":{"query":"gmail"} + }] + }`), &response)) + require.Len(t, response.Output, 1) + require.Equal(t, `{"query":"gmail"}`, response.Output[0].Arguments) +} + +func TestResponsesStreamEventUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var event ResponsesStreamEvent + require.NoError(t, json.Unmarshal([]byte(`{ + "type":"response.output_item.done", + "item":{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "arguments":{"query":"gmail"} + } + }`), &event)) + require.NotNil(t, event.Item) + require.Equal(t, `{"query":"gmail"}`, event.Item.Arguments) +} diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic.go b/backend/internal/pkg/apicompat/responses_to_anthropic.go index 037e16b652..376f0d97da 100644 --- a/backend/internal/pkg/apicompat/responses_to_anthropic.go +++ b/backend/internal/pkg/apicompat/responses_to_anthropic.go @@ -100,15 +100,16 @@ func anthropicUsageFromResponsesUsage(usage *ResponsesUsage) AnthropicUsage { cachedTokens = usage.InputTokensDetails.CachedTokens } - inputTokens := usage.InputTokens - cachedTokens + inputTokens := usage.InputTokens - cachedTokens - usage.CacheCreationInputTokens if inputTokens < 0 { inputTokens = 0 } return AnthropicUsage{ - InputTokens: inputTokens, - OutputTokens: usage.OutputTokens, - CacheReadInputTokens: cachedTokens, + InputTokens: inputTokens, + OutputTokens: usage.OutputTokens, + CacheReadInputTokens: cachedTokens, + CacheCreationInputTokens: usage.CacheCreationInputTokens, } } @@ -181,9 +182,10 @@ type ResponsesEventToAnthropicState struct { // OutputIndexToBlockIdx maps Responses output_index → Anthropic content block index. OutputIndexToBlockIdx map[int]int - InputTokens int - OutputTokens int - CacheReadInputTokens int + InputTokens int + OutputTokens int + CacheReadInputTokens int + CacheCreationInputTokens int ResponseID string Model string @@ -258,9 +260,10 @@ func FinalizeResponsesAnthropicStream(state *ResponsesEventToAnthropicState) []A StopReason: stopReason, }, Usage: &AnthropicUsage{ - InputTokens: state.InputTokens, - OutputTokens: state.OutputTokens, - CacheReadInputTokens: state.CacheReadInputTokens, + InputTokens: state.InputTokens, + OutputTokens: state.OutputTokens, + CacheReadInputTokens: state.CacheReadInputTokens, + CacheCreationInputTokens: state.CacheCreationInputTokens, }, }, AnthropicStreamEvent{Type: "message_stop"}, @@ -410,10 +413,6 @@ func resToAnthHandleFuncArgsDelta(evt *ResponsesStreamEvent, state *ResponsesEve return nil } - if state.CurrentBlockType == "tool_use" && state.CurrentToolName == "Read" { - state.CurrentToolArgs += evt.Delta - return nil - } if state.CurrentBlockType == "tool_use" { state.CurrentToolHadDelta = true } @@ -578,6 +577,7 @@ func resToAnthHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo state.InputTokens = usage.InputTokens state.OutputTokens = usage.OutputTokens state.CacheReadInputTokens = usage.CacheReadInputTokens + state.CacheCreationInputTokens = usage.CacheCreationInputTokens } if evt.Response != nil { if evt.Response.Usage != nil { @@ -585,6 +585,7 @@ func resToAnthHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo state.InputTokens = usage.InputTokens state.OutputTokens = usage.OutputTokens state.CacheReadInputTokens = usage.CacheReadInputTokens + state.CacheCreationInputTokens = usage.CacheCreationInputTokens } switch evt.Response.Status { case "incomplete": @@ -605,9 +606,10 @@ func resToAnthHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo StopReason: stopReason, }, Usage: &AnthropicUsage{ - InputTokens: state.InputTokens, - OutputTokens: state.OutputTokens, - CacheReadInputTokens: state.CacheReadInputTokens, + InputTokens: state.InputTokens, + OutputTokens: state.OutputTokens, + CacheReadInputTokens: state.CacheReadInputTokens, + CacheCreationInputTokens: state.CacheCreationInputTokens, }, }, AnthropicStreamEvent{Type: "message_stop"}, diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_read_tool_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_read_tool_test.go new file mode 100644 index 0000000000..72b60099fe --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_to_anthropic_read_tool_test.go @@ -0,0 +1,84 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestResToAnthFuncArgsDelta_ReadToolStreamsDeltas(t *testing.T) { + state := NewResponsesEventToAnthropicState() + state.MessageStartSent = true + state.CurrentBlockType = "tool_use" + state.CurrentToolName = "Read" + state.OutputIndexToBlockIdx = map[int]int{0: 0} + + evt := &ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 0, + Delta: `{"file_path":"/tmp/test.go"}`, + } + + events := ResponsesEventToAnthropicEvents(evt, state) + + require.Len(t, events, 1, "Read tool delta must produce content_block_delta") + assert.Equal(t, "content_block_delta", events[0].Type) + assert.Equal(t, "input_json_delta", events[0].Delta.Type) + assert.Equal(t, `{"file_path":"/tmp/test.go"}`, events[0].Delta.PartialJSON) + assert.True(t, state.CurrentToolHadDelta, "Read deltas should set CurrentToolHadDelta") +} + +func TestResToAnthFuncArgsDelta_ReadToolWithoutDone(t *testing.T) { + state := NewResponsesEventToAnthropicState() + state.MessageStartSent = true + state.ContentBlockIndex = 0 + state.ContentBlockOpen = true + state.CurrentBlockType = "tool_use" + state.CurrentToolName = "Read" + state.OutputIndexToBlockIdx = map[int]int{0: 0} + + delta := &ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 0, + Delta: `{"file_path":"/tmp/test.go"}`, + } + events := ResponsesEventToAnthropicEvents(delta, state) + require.Len(t, events, 1, "delta should be streamed") + + completed := &ResponsesStreamEvent{ + Type: "response.completed", + Response: &ResponsesResponse{ + Status: "completed", + }, + } + events = ResponsesEventToAnthropicEvents(completed, state) + + hasStop := false + for _, e := range events { + if e.Type == "content_block_stop" { + hasStop = true + } + } + assert.True(t, hasStop, "block should be closed even without .done event") +} + +func TestResToAnthFuncArgsDelta_NonReadToolUnchanged(t *testing.T) { + state := NewResponsesEventToAnthropicState() + state.MessageStartSent = true + state.CurrentBlockType = "tool_use" + state.CurrentToolName = "Write" + state.OutputIndexToBlockIdx = map[int]int{0: 0} + + evt := &ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 0, + Delta: `{"file_path":"/tmp/out.txt","content":"hello"}`, + } + + events := ResponsesEventToAnthropicEvents(evt, state) + + require.Len(t, events, 1) + assert.Equal(t, "content_block_delta", events[0].Type) + assert.True(t, state.CurrentToolHadDelta) +} diff --git a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go index 2ae6f8ac3f..a89a1b4203 100644 --- a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go +++ b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go @@ -89,8 +89,13 @@ func ResponsesToChatCompletions(resp *ResponsesResponse, model string) *ChatComp func responsesStatusToChatFinishReason(status string, details *ResponsesIncompleteDetails, toolCalls []ChatToolCall) string { switch status { case "incomplete": - if details != nil && details.Reason == "max_output_tokens" { - return "length" + if details != nil { + switch details.Reason { + case "max_output_tokens": + return "length" + case "content_filter": + return "content_filter" + } } return "stop" case "completed": @@ -299,8 +304,13 @@ func resToChatHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo switch evt.Response.Status { case "incomplete": - if evt.Response.IncompleteDetails != nil && evt.Response.IncompleteDetails.Reason == "max_output_tokens" { - finishReason = "length" + if evt.Response.IncompleteDetails != nil { + switch evt.Response.IncompleteDetails.Reason { + case "max_output_tokens": + finishReason = "length" + case "content_filter": + finishReason = "content_filter" + } } case "completed": if state.SawToolCall { diff --git a/backend/internal/pkg/apicompat/streaming_stop_reason_test.go b/backend/internal/pkg/apicompat/streaming_stop_reason_test.go new file mode 100644 index 0000000000..c2889f0251 --- /dev/null +++ b/backend/internal/pkg/apicompat/streaming_stop_reason_test.go @@ -0,0 +1,122 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAnthropicStreamingMaxTokens_MapsToIncomplete(t *testing.T) { + state := NewAnthropicEventToResponsesState() + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_start", + Message: &AnthropicResponse{ID: "msg_test", Model: "claude-opus-4-6", Role: "assistant"}, + }, state) + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_delta", + Delta: &AnthropicDelta{ + StopReason: "max_tokens", + }, + Usage: &AnthropicUsage{OutputTokens: 4096}, + }, state) + + require.Equal(t, "max_tokens", state.StopReason) + + events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_stop", + }, state) + + var completed *ResponsesStreamEvent + for i := range events { + if events[i].Type == "response.completed" || events[i].Type == "response.incomplete" { + completed = &events[i] + break + } + } + require.NotNil(t, completed, "should have terminal event") + assert.Equal(t, "response.incomplete", completed.Type) + require.NotNil(t, completed.Response) + assert.Equal(t, "incomplete", completed.Response.Status) + require.NotNil(t, completed.Response.IncompleteDetails) + assert.Equal(t, "max_output_tokens", completed.Response.IncompleteDetails.Reason) +} + +func TestAnthropicStreamingEndTurn_MapsToCompleted(t *testing.T) { + state := NewAnthropicEventToResponsesState() + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_start", + Message: &AnthropicResponse{ID: "msg_test", Model: "claude-opus-4-6", Role: "assistant"}, + }, state) + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_delta", + Delta: &AnthropicDelta{StopReason: "end_turn"}, + Usage: &AnthropicUsage{OutputTokens: 100}, + }, state) + + events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_stop", + }, state) + + var completed *ResponsesStreamEvent + for i := range events { + if events[i].Type == "response.completed" { + completed = &events[i] + break + } + } + require.NotNil(t, completed) + assert.Equal(t, "completed", completed.Response.Status) + assert.Nil(t, completed.Response.IncompleteDetails) +} + +func TestResponsesToChatCompletions_ContentFilter(t *testing.T) { + resp := &ResponsesResponse{ + ID: "resp_cf", + Status: "incomplete", + IncompleteDetails: &ResponsesIncompleteDetails{ + Reason: "content_filter", + }, + Output: []ResponsesOutput{{ + Type: "message", + Content: []ResponsesContentPart{{Type: "output_text", Text: "partial"}}, + }}, + Usage: &ResponsesUsage{InputTokens: 10, OutputTokens: 5}, + } + + cc := ResponsesToChatCompletions(resp, "gpt-5.5") + require.Len(t, cc.Choices, 1) + assert.Equal(t, "content_filter", cc.Choices[0].FinishReason) +} + +func TestResponsesToChatCompletionsStreaming_ContentFilter(t *testing.T) { + state := NewResponsesEventToChatState() + state.ID = "resp_cf" + state.Model = "gpt-5.5" + state.SentRole = true + + events := ResponsesEventToChatChunks(&ResponsesStreamEvent{ + Type: "response.completed", + Response: &ResponsesResponse{ + ID: "resp_cf", + Status: "incomplete", + IncompleteDetails: &ResponsesIncompleteDetails{ + Reason: "content_filter", + }, + }, + }, state) + + hasContentFilter := false + for _, chunk := range events { + for _, choice := range chunk.Choices { + if choice.FinishReason != nil && *choice.FinishReason == "content_filter" { + hasContentFilter = true + } + } + } + assert.True(t, hasContentFilter, "streaming content_filter should map to finish_reason content_filter") +} diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index 8d96a1d3d3..6cf9a2be31 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -249,11 +249,31 @@ type ResponsesContentPart struct { // ResponsesTool describes a tool in the Responses API. type ResponsesTool struct { - Type string `json:"type"` // "function" | "web_search" | "local_shell" etc. + Type string `json:"type"` // "function" | "custom" | "web_search" | "local_shell" etc. Name string `json:"name,omitempty"` Description string `json:"description,omitempty"` Parameters json.RawMessage `json:"parameters,omitempty"` Strict *bool `json:"strict,omitempty"` + + // type=namespace 的子工具列表(tools 与 children 二选一,语义相同)。 + Tools []ResponsesTool `json:"tools,omitempty"` + Children []ResponsesTool `json:"children,omitempty"` +} + +// UnmarshalJSON 容忍字符串形式的工具声明:codex 会以 "name" 简写声明 custom 工具, +func (t *ResponsesTool) UnmarshalJSON(data []byte) error { + var name string + if err := json.Unmarshal(data, &name); err == nil { + *t = ResponsesTool{Type: "custom", Name: name} + return nil + } + type alias ResponsesTool + var a alias + if err := json.Unmarshal(data, &a); err != nil { + return err + } + *t = ResponsesTool(a) + return nil } // ResponsesResponse is the non-streaming response from POST /v1/responses. @@ -301,11 +321,88 @@ type ResponsesOutput struct { CallID string `json:"call_id,omitempty"` Name string `json:"name,omitempty"` Arguments string `json:"arguments,omitempty"` + // 来源为 namespace 子工具时的归属命名空间(codex 按 namespace+name 路由该调用)。 + Namespace string `json:"namespace,omitempty"` + + // type=custom_tool_call(custom/freeform 工具,input 为自由文本) + Input string `json:"input,omitempty"` // type=web_search_call Action *WebSearchAction `json:"action,omitempty"` } +// MarshalJSON 处理 tool_search_call 项的线上形态(复用 CallID/Arguments 字段): +// execution 固定为 "client"(codex 的必填字段,非 client 的调用会被静默忽略), +// arguments 是 JSON 对象而非 function_call 语义下的字符串。其余类型走默认结构体 +// 序列化,输出逐字节不变。 +func (o ResponsesOutput) MarshalJSON() ([]byte, error) { + type responsesOutputAlias ResponsesOutput + if o.Type != "tool_search_call" { + return json.Marshal(responsesOutputAlias(o)) + } + m := map[string]any{ + "type": o.Type, + "id": o.ID, + "call_id": o.CallID, + "execution": "client", + "arguments": toolSearchCallArgumentsJSON(o.Arguments), + } + if o.Status != "" { + m["status"] = o.Status + } + return json.Marshal(m) +} + +// UnmarshalJSON accepts both the Responses function-call string form and the +// tool_search_call object form for arguments. The bridge stores arguments as a +// string internally, so object arguments are retained as their raw JSON. +func (o *ResponsesOutput) UnmarshalJSON(data []byte) error { + type responsesOutputAlias ResponsesOutput + + var kind struct { + Type string `json:"type"` + } + if err := json.Unmarshal(data, &kind); err != nil { + return err + } + if kind.Type != "tool_search_call" { + var decoded responsesOutputAlias + if err := json.Unmarshal(data, &decoded); err != nil { + return err + } + *o = ResponsesOutput(decoded) + return nil + } + + var fields map[string]json.RawMessage + if err := json.Unmarshal(data, &fields); err != nil { + return err + } + arguments, hasArguments := fields["arguments"] + delete(fields, "arguments") + normalized, err := json.Marshal(fields) + if err != nil { + return err + } + + var decoded responsesOutputAlias + if err := json.Unmarshal(normalized, &decoded); err != nil { + return err + } + *o = ResponsesOutput(decoded) + if !hasArguments || string(arguments) == "null" { + return nil + } + + var argumentString string + if err := json.Unmarshal(arguments, &argumentString); err == nil { + o.Arguments = argumentString + } else { + o.Arguments = string(arguments) + } + return nil +} + // WebSearchAction describes the search action in a web_search_call output item. type WebSearchAction struct { Type string `json:"type,omitempty"` // "search" @@ -444,6 +541,9 @@ type ResponsesStreamEvent struct { Name string `json:"name,omitempty"` Arguments string `json:"arguments,omitempty"` + // response.custom_tool_call_input.done + Input string `json:"input,omitempty"` + // response.reasoning_summary_text.delta / done // Reuses Text/Delta fields above, SummaryIndex identifies which summary part SummaryIndex int `json:"summary_index,omitempty"` diff --git a/backend/internal/pkg/httpclient/pool.go b/backend/internal/pkg/httpclient/pool.go index 12804cc67d..22d3c65feb 100644 --- a/backend/internal/pkg/httpclient/pool.go +++ b/backend/internal/pkg/httpclient/pool.go @@ -25,6 +25,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" ) @@ -92,6 +93,7 @@ func buildClient(opts Options) (*http.Client, error) { if opts.ValidateResolvedIP && !opts.AllowPrivateHosts { rt = newValidatedTransport(transport) } + rt = servertiming.WrapRoundTripper(rt) return &http.Client{ Transport: rt, Timeout: opts.Timeout, diff --git a/backend/internal/pkg/pagination/pagination.go b/backend/internal/pkg/pagination/pagination.go index ce8e74b8ce..334ba809de 100644 --- a/backend/internal/pkg/pagination/pagination.go +++ b/backend/internal/pkg/pagination/pagination.go @@ -38,7 +38,7 @@ func (p PaginationParams) Offset() int { if p.Page < 1 { p.Page = 1 } - return (p.Page - 1) * p.PageSize + return (p.Page - 1) * p.Limit() } // Limit 获取限制数 diff --git a/backend/internal/pkg/pagination/pagination_test.go b/backend/internal/pkg/pagination/pagination_test.go index 9a3b069d90..9704449e92 100644 --- a/backend/internal/pkg/pagination/pagination_test.go +++ b/backend/internal/pkg/pagination/pagination_test.go @@ -69,3 +69,30 @@ func TestPaginationParamsLimit(t *testing.T) { }) } } + +func TestPaginationParamsOffsetUsesNormalizedLimit(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + page int + pageSize int + want int + }{ + {name: "invalid page uses first page", page: 0, pageSize: 50, want: 0}, + {name: "zero page size uses default", page: 2, pageSize: 0, want: 20}, + {name: "negative page size uses default", page: 2, pageSize: -1, want: 20}, + {name: "normal values", page: 3, pageSize: 50, want: 100}, + {name: "page size beyond max is clamped", page: 2, pageSize: 1500, want: 1000}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + params := PaginationParams{Page: tt.page, PageSize: tt.pageSize} + if got := params.Offset(); got != tt.want { + t.Fatalf("Offset() for Page=%d, PageSize=%d = %d, want %d", tt.page, tt.pageSize, got, tt.want) + } + }) + } +} diff --git a/backend/internal/pkg/servertiming/collector.go b/backend/internal/pkg/servertiming/collector.go new file mode 100644 index 0000000000..553edede31 --- /dev/null +++ b/backend/internal/pkg/servertiming/collector.go @@ -0,0 +1,348 @@ +package servertiming + +import ( + "context" + "fmt" + "sort" + "strconv" + "strings" + "sync" + "time" +) + +const ( + HeaderName = "Server-Timing" + AdminUIHeader = "X-Admin-UI-Request" + MetricDatabase = "db" + MetricRedis = "redis" + dependencyPrefix = "dep_" + + maxMetricNameLength = 48 + maxIntervals = 2048 + maxHeaderLength = 4096 +) + +type contextKey struct{} + +type interval struct { + start time.Time + end time.Time +} + +type metric struct { + count int64 + intervals []interval +} + +// Collector stores request-scoped timing samples. It is safe for concurrent use. +type Collector struct { + startedAt time.Time + + mu sync.Mutex + metrics map[string]*metric + cacheStatus string +} + +// New creates a collector whose total duration starts at startedAt. +func New(startedAt time.Time) *Collector { + if startedAt.IsZero() { + startedAt = time.Now() + } + return &Collector{ + startedAt: startedAt, + metrics: make(map[string]*metric), + } +} + +// WithCollector attaches a collector to a context. +func WithCollector(ctx context.Context, collector *Collector) context.Context { + if ctx == nil { + ctx = context.Background() + } + if collector == nil { + return ctx + } + return context.WithValue(ctx, contextKey{}, collector) +} + +// FromContext returns the request timing collector, when one is active. +func FromContext(ctx context.Context) (*Collector, bool) { + if ctx == nil { + return nil, false + } + collector, ok := ctx.Value(contextKey{}).(*Collector) + return collector, ok && collector != nil +} + +// Active reports whether timing collection is enabled for this request. +func Active(ctx context.Context) bool { + _, ok := FromContext(ctx) + return ok +} + +// Record adds a completed interval and operation count to a metric. +func Record(ctx context.Context, name string, startedAt, endedAt time.Time, count int) { + collector, ok := FromContext(ctx) + if !ok { + return + } + collector.Record(name, startedAt, endedAt, count) +} + +// RecordInterval adds timing without incrementing the operation count. It is +// useful when one logical operation has multiple blocking driver calls. +func RecordInterval(ctx context.Context, name string, startedAt, endedAt time.Time) { + collector, ok := FromContext(ctx) + if !ok { + return + } + collector.record(name, startedAt, endedAt, 0) +} + +// Record adds a completed interval directly to the collector. +func (c *Collector) Record(name string, startedAt, endedAt time.Time, count int) { + if count <= 0 { + count = 1 + } + c.record(name, startedAt, endedAt, count) +} + +func (c *Collector) record(name string, startedAt, endedAt time.Time, count int) { + name = normalizeMetricName(name) + if c == nil || name == "" || startedAt.IsZero() || endedAt.Before(startedAt) { + return + } + if count < 0 { + count = 0 + } + + c.mu.Lock() + m := c.metrics[name] + if m == nil { + m = &metric{} + c.metrics[name] = m + } + m.count += int64(count) + if len(m.intervals) < maxIntervals { + m.intervals = append(m.intervals, interval{start: startedAt, end: endedAt}) + } + c.mu.Unlock() +} + +// Observe starts a metric span and returns an idempotent completion function. +func Observe(ctx context.Context, name string) func() { + collector, ok := FromContext(ctx) + name = normalizeMetricName(name) + if !ok || name == "" { + return func() {} + } + startedAt := time.Now() + var once sync.Once + return func() { + once.Do(func() { + collector.Record(name, startedAt, time.Now(), 1) + }) + } +} + +// ObserveDependency starts a named external dependency span. +func ObserveDependency(ctx context.Context, module string) func() { + return Observe(ctx, dependencyMetricName(module)) +} + +// RecordDependency records a completed external dependency interval. +func RecordDependency(ctx context.Context, module string, startedAt, endedAt time.Time) { + Record(ctx, dependencyMetricName(module), startedAt, endedAt, 1) +} + +// SetCacheStatus records the response-cache outcome for the request. +func SetCacheStatus(ctx context.Context, status string) { + collector, ok := FromContext(ctx) + if !ok { + return + } + status = normalizeCacheStatus(status) + if status == "" { + return + } + collector.mu.Lock() + collector.cacheStatus = status + collector.mu.Unlock() +} + +// HeaderValue renders a bounded, deterministic Server-Timing header. +func HeaderValue(ctx context.Context, endedAt time.Time, cacheStatus string) string { + collector, ok := FromContext(ctx) + if !ok { + return "" + } + return collector.HeaderValue(endedAt, cacheStatus) +} + +// HeaderValue renders a bounded, deterministic Server-Timing header. +func (c *Collector) HeaderValue(endedAt time.Time, cacheStatus string) string { + if c == nil { + return "" + } + if endedAt.IsZero() { + endedAt = time.Now() + } + if endedAt.Before(c.startedAt) { + endedAt = c.startedAt + } + + c.mu.Lock() + metrics := make(map[string]metric, len(c.metrics)) + allIntervals := make([]interval, 0) + dependencyIntervals := make([]interval, 0) + var dependencyCount int64 + for name, source := range c.metrics { + copied := metric{count: source.count, intervals: append([]interval(nil), source.intervals...)} + metrics[name] = copied + allIntervals = append(allIntervals, copied.intervals...) + if strings.HasPrefix(name, dependencyPrefix) { + dependencyIntervals = append(dependencyIntervals, copied.intervals...) + dependencyCount += copied.count + } + } + storedCacheStatus := c.cacheStatus + c.mu.Unlock() + + total := endedAt.Sub(c.startedAt) + blocked := unionDuration(allIntervals, c.startedAt, endedAt) + app := total - blocked + if app < 0 { + app = 0 + } + + cacheStatus = normalizeCacheStatus(cacheStatus) + if cacheStatus == "" { + cacheStatus = normalizeCacheStatus(storedCacheStatus) + } + if cacheStatus == "" { + cacheStatus = "bypass" + } + + database := metrics[MetricDatabase] + redisMetric := metrics[MetricRedis] + parts := []string{ + "total;dur=" + formatDuration(total), + "app;dur=" + formatDuration(app), + fmt.Sprintf("db;dur=%s;desc=\"queries=%d\"", formatDuration(unionDuration(database.intervals, c.startedAt, endedAt)), database.count), + fmt.Sprintf("redis;dur=%s;desc=\"commands=%d\"", formatDuration(unionDuration(redisMetric.intervals, c.startedAt, endedAt)), redisMetric.count), + "cache;desc=\"" + cacheStatus + "\"", + fmt.Sprintf("deps;dur=%s;desc=\"calls=%d\"", formatDuration(unionDuration(dependencyIntervals, c.startedAt, endedAt)), dependencyCount), + } + + dependencyNames := make([]string, 0) + for name := range metrics { + if strings.HasPrefix(name, dependencyPrefix) { + dependencyNames = append(dependencyNames, name) + } + } + sort.Strings(dependencyNames) + for _, name := range dependencyNames { + m := metrics[name] + part := fmt.Sprintf("%s;dur=%s;desc=\"calls=%d\"", name, formatDuration(unionDuration(m.intervals, c.startedAt, endedAt)), m.count) + candidate := strings.Join(append(parts, part), ", ") + if len(candidate) > maxHeaderLength { + break + } + parts = append(parts, part) + } + + return strings.Join(parts, ", ") +} + +func dependencyMetricName(module string) string { + module = normalizeMetricName(module) + module = strings.TrimPrefix(module, dependencyPrefix) + if module == "" { + module = "http" + } + return dependencyPrefix + module +} + +func normalizeMetricName(name string) string { + name = strings.ToLower(strings.TrimSpace(name)) + if name == "" { + return "" + } + var b strings.Builder + b.Grow(min(len(name), maxMetricNameLength)) + for _, r := range name { + if b.Len() >= maxMetricNameLength { + break + } + switch { + case r >= 'a' && r <= 'z', r >= '0' && r <= '9': + _, _ = b.WriteRune(r) + case r == '_' || r == '-': + _ = b.WriteByte('_') + } + } + return strings.Trim(b.String(), "_") +} + +func normalizeCacheStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case "hit": + return "hit" + case "miss": + return "miss" + case "bypass": + return "bypass" + default: + return "" + } +} + +func unionDuration(intervals []interval, lowerBound, upperBound time.Time) time.Duration { + if len(intervals) == 0 || !upperBound.After(lowerBound) { + return 0 + } + normalized := make([]interval, 0, len(intervals)) + for _, item := range intervals { + start := item.start + end := item.end + if start.Before(lowerBound) { + start = lowerBound + } + if end.After(upperBound) { + end = upperBound + } + if end.After(start) { + normalized = append(normalized, interval{start: start, end: end}) + } + } + if len(normalized) == 0 { + return 0 + } + sort.Slice(normalized, func(i, j int) bool { + return normalized[i].start.Before(normalized[j].start) + }) + + currentStart := normalized[0].start + currentEnd := normalized[0].end + var total time.Duration + for _, item := range normalized[1:] { + if !item.start.After(currentEnd) { + if item.end.After(currentEnd) { + currentEnd = item.end + } + continue + } + total += currentEnd.Sub(currentStart) + currentStart = item.start + currentEnd = item.end + } + total += currentEnd.Sub(currentStart) + return total +} + +func formatDuration(value time.Duration) string { + if value < 0 { + value = 0 + } + return strconv.FormatFloat(float64(value)/float64(time.Millisecond), 'f', 1, 64) +} diff --git a/backend/internal/pkg/servertiming/collector_test.go b/backend/internal/pkg/servertiming/collector_test.go new file mode 100644 index 0000000000..1bb809f9bc --- /dev/null +++ b/backend/internal/pkg/servertiming/collector_test.go @@ -0,0 +1,129 @@ +package servertiming + +import ( + "context" + "fmt" + "strings" + "sync" + "testing" + "time" +) + +func TestCollectorHeaderValueAggregatesIntervals(t *testing.T) { + startedAt := time.Unix(100, 0) + collector := New(startedAt) + collector.Record(MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(40*time.Millisecond), 2) + collector.Record(MetricRedis, startedAt.Add(30*time.Millisecond), startedAt.Add(50*time.Millisecond), 3) + collector.Record(dependencyMetricName("openai"), startedAt.Add(70*time.Millisecond), startedAt.Add(100*time.Millisecond), 1) + collector.Record(dependencyMetricName("github"), startedAt.Add(60*time.Millisecond), startedAt.Add(90*time.Millisecond), 1) + + got := collector.HeaderValue(startedAt.Add(120*time.Millisecond), "miss") + want := `total;dur=120.0, app;dur=40.0, db;dur=30.0;desc="queries=2", redis;dur=20.0;desc="commands=3", cache;desc="miss", deps;dur=40.0;desc="calls=2", dep_github;dur=30.0;desc="calls=1", dep_openai;dur=30.0;desc="calls=1"` + if got != want { + t.Fatalf("HeaderValue() = %q, want %q", got, want) + } +} + +func TestRecordIntervalDoesNotIncrementCount(t *testing.T) { + startedAt := time.Unix(200, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + Record(ctx, MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(20*time.Millisecond), 1) + RecordInterval(ctx, MetricDatabase, startedAt.Add(30*time.Millisecond), startedAt.Add(40*time.Millisecond)) + + header := HeaderValue(ctx, startedAt.Add(100*time.Millisecond), "hit") + if !strings.Contains(header, `db;dur=20.0;desc="queries=1"`) { + t.Fatalf("header %q does not contain one query with both blocking intervals", header) + } + if !strings.Contains(header, "app;dur=80.0") { + t.Fatalf("header %q does not subtract the interval union from app time", header) + } +} + +func TestCollectorCacheStatusFallback(t *testing.T) { + startedAt := time.Unix(300, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + SetCacheStatus(ctx, " HIT ") + if got := HeaderValue(ctx, startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="hit"`) { + t.Fatalf("HeaderValue() = %q, want stored cache hit", got) + } + + other := New(startedAt) + if got := other.HeaderValue(startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="bypass"`) { + t.Fatalf("HeaderValue() = %q, want cache bypass", got) + } +} + +func TestCollectorSanitizesDependencyMetric(t *testing.T) { + startedAt := time.Unix(400, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + RecordDependency(ctx, "GitHub API\r\nInjected;dur=999", startedAt, startedAt.Add(time.Millisecond)) + + header := HeaderValue(ctx, startedAt.Add(2*time.Millisecond), "bypass") + if strings.ContainsAny(header, "\r\n") || strings.Contains(header, ";dur=999") { + t.Fatalf("unsafe metric content reached header: %q", header) + } + if !strings.Contains(header, "dep_githubapiinjecteddur999;dur=1.0") { + t.Fatalf("sanitized dependency metric missing from header: %q", header) + } +} + +func TestCollectorBoundsHeaderLength(t *testing.T) { + startedAt := time.Unix(500, 0) + collector := New(startedAt) + for i := 0; i < 300; i++ { + collector.Record( + dependencyMetricName(fmt.Sprintf("module_%03d_with_a_deliberately_long_name", i)), + startedAt, + startedAt.Add(time.Millisecond), + 1, + ) + } + + header := collector.HeaderValue(startedAt.Add(2*time.Millisecond), "bypass") + if len(header) > maxHeaderLength { + t.Fatalf("header length = %d, want <= %d", len(header), maxHeaderLength) + } + if !strings.Contains(header, "total;dur=2.0") || !strings.Contains(header, "deps;dur=1.0") { + t.Fatalf("bounded header lost fixed metrics: %q", header) + } +} + +func TestCollectorConcurrentRecording(t *testing.T) { + startedAt := time.Now() + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + const workers = 25 + const recordsPerWorker = 100 + var wg sync.WaitGroup + wg.Add(workers) + for i := 0; i < workers; i++ { + go func() { + defer wg.Done() + for j := 0; j < recordsPerWorker; j++ { + Record(ctx, MetricDatabase, startedAt, startedAt.Add(time.Microsecond), 1) + } + }() + } + wg.Wait() + + header := HeaderValue(ctx, startedAt.Add(time.Millisecond), "bypass") + want := fmt.Sprintf(`queries=%d`, workers*recordsPerWorker) + if !strings.Contains(header, want) { + t.Fatalf("header %q does not contain %q", header, want) + } +} + +func TestContextHelpersHandleMissingCollector(t *testing.T) { + if Active(context.Background()) { + t.Fatal("context without collector reported active") + } + if got := HeaderValue(context.Background(), time.Now(), "hit"); got != "" { + t.Fatalf("HeaderValue() = %q without collector, want empty", got) + } +} diff --git a/backend/internal/pkg/servertiming/http.go b/backend/internal/pkg/servertiming/http.go new file mode 100644 index 0000000000..e326e24302 --- /dev/null +++ b/backend/internal/pkg/servertiming/http.go @@ -0,0 +1,104 @@ +package servertiming + +import ( + "context" + "net/http" + "strings" + "time" +) + +type dependencyModuleKey struct{} + +type timingRoundTripper struct { + base http.RoundTripper +} + +// WithDependencyModule overrides the safe module name used for an outbound call. +func WithDependencyModule(ctx context.Context, module string) context.Context { + if ctx == nil { + ctx = context.Background() + } + module = strings.TrimPrefix(normalizeMetricName(module), dependencyPrefix) + if module == "" { + return ctx + } + return context.WithValue(ctx, dependencyModuleKey{}, module) +} + +// WrapRoundTripper records outbound response-header latency for active requests. +func WrapRoundTripper(base http.RoundTripper) http.RoundTripper { + if base == nil { + base = http.DefaultTransport + } + if _, ok := base.(*timingRoundTripper); ok { + return base + } + return &timingRoundTripper{base: base} +} + +// InstrumentClient returns a shallow client copy with an instrumented transport. +func InstrumentClient(client *http.Client) *http.Client { + if client == nil { + client = &http.Client{} + } + copyClient := *client + copyClient.Transport = WrapRoundTripper(copyClient.Transport) + return ©Client +} + +// Do records response-header latency without changing the client's transport +// type. Use it for clients whose callers inspect or configure *http.Transport. +func Do(client *http.Client, req *http.Request) (*http.Response, error) { + if client == nil { + client = http.DefaultClient + } + if req == nil || !Active(req.Context()) { + return client.Do(req) + } + startedAt := time.Now() + response, err := client.Do(req) + RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now()) + return response, err +} + +func (t *timingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if req == nil || !Active(req.Context()) { + return t.base.RoundTrip(req) + } + startedAt := time.Now() + response, err := t.base.RoundTrip(req) + RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now()) + return response, err +} + +func dependencyModule(req *http.Request) string { + if req != nil { + if module, ok := req.Context().Value(dependencyModuleKey{}).(string); ok && module != "" { + return module + } + } + if req == nil || req.URL == nil { + return "http" + } + host := strings.ToLower(req.URL.Hostname()) + switch { + case strings.Contains(host, "github"): + return "github" + case strings.Contains(host, "openai"): + return "openai" + case strings.Contains(host, "anthropic"): + return "anthropic" + case strings.Contains(host, "generativelanguage") || strings.Contains(host, "gemini"): + return "gemini" + case strings.Contains(host, "cloudcode") || strings.Contains(host, "antigravity"): + return "antigravity" + case strings.Contains(host, "googleapis") || strings.Contains(host, "google"): + return "google" + case strings.Contains(host, "amazonaws") || strings.Contains(host, "cloudflarestorage") || strings.Contains(host, "s3"): + return "s3" + case strings.Contains(host, "stripe") || strings.Contains(host, "airwallex") || strings.Contains(host, "alipay") || strings.Contains(host, "wechatpay") || strings.Contains(host, "paypal"): + return "payment" + default: + return "http" + } +} diff --git a/backend/internal/pkg/servertiming/http_test.go b/backend/internal/pkg/servertiming/http_test.go new file mode 100644 index 0000000000..d37f378414 --- /dev/null +++ b/backend/internal/pkg/servertiming/http_test.go @@ -0,0 +1,168 @@ +package servertiming + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + "time" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +type trackingBody struct { + read bool +} + +func (b *trackingBody) Read(_ []byte) (int, error) { + b.read = true + return 0, io.EOF +} + +func (b *trackingBody) Close() error { return nil } + +func TestWrapRoundTripperRecordsResponseHeaderLatency(t *testing.T) { + startedAt := time.Now() + collector := New(startedAt) + body := &trackingBody{} + baseCalled := false + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + baseCalled = true + return &http.Response{ + StatusCode: http.StatusOK, + Body: body, + Header: make(http.Header), + Request: req, + }, nil + }) + req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.github.com/repos/example/project", nil) + if err != nil { + t.Fatal(err) + } + + resp, err := WrapRoundTripper(base).RoundTrip(req) + if err != nil { + t.Fatal(err) + } + defer func() { _ = resp.Body.Close() }() + if !baseCalled { + t.Fatal("base RoundTripper was not called") + } + if body.read { + t.Fatal("RoundTripper instrumentation read the response body; timing must stop at response headers") + } + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `dep_github;dur=`) || !strings.Contains(header, `deps;dur=`) { + t.Fatalf("dependency metrics missing from header: %q", header) + } +} + +func TestWrapRoundTripperUsesContextModuleOverride(t *testing.T) { + collector := New(time.Now()) + ctx := WithDependencyModule(WithCollector(context.Background(), collector), "data-managementd") + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://private.example.test/path", nil) + if err != nil { + t.Fatal(err) + } + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + + if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil { + t.Fatal(err) + } + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, "dep_data_managementd") { + t.Fatalf("module override missing from header: %q", header) + } + if strings.Contains(header, "private.example") { + t.Fatalf("raw host leaked into header: %q", header) + } +} + +func TestWrapRoundTripperSkipsInactiveContext(t *testing.T) { + called := false + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + called = true + return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + req, err := http.NewRequest(http.MethodGet, "https://api.openai.com/v1/models", nil) + if err != nil { + t.Fatal(err) + } + if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil { + t.Fatal(err) + } + if !called { + t.Fatal("inactive request did not reach base RoundTripper") + } +} + +func TestDoRecordsWithoutChangingTransportType(t *testing.T) { + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + client := &http.Client{Transport: base} + collector := New(time.Now()) + req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.openai.com/v1/models", nil) + if err != nil { + t.Fatal(err) + } + if _, err := Do(client, req); err != nil { + t.Fatal(err) + } + if _, ok := client.Transport.(roundTripFunc); !ok { + t.Fatalf("Do changed client transport type to %T", client.Transport) + } + if header := collector.HeaderValue(time.Now(), "bypass"); !strings.Contains(header, "dep_openai;dur=") { + t.Fatalf("dependency metric missing from header: %q", header) + } +} + +func TestDependencyModuleClassification(t *testing.T) { + tests := map[string]string{ + "https://api.github.com/repos/a/b": "github", + "https://api.openai.com/v1/models": "openai", + "https://api.anthropic.com/v1/messages": "anthropic", + "https://generativelanguage.googleapis.com/v1/models": "gemini", + "https://cloudcode-pa.googleapis.com/v1internal": "antigravity", + "https://storage.googleapis.com/bucket/object": "google", + "https://bucket.s3.amazonaws.com/object": "s3", + "https://api.stripe.com/v1/refunds": "payment", + "https://dependency.example.test/path": "http", + } + for rawURL, want := range tests { + req, err := http.NewRequest(http.MethodGet, rawURL, nil) + if err != nil { + t.Fatalf("NewRequest(%q): %v", rawURL, err) + } + if got := dependencyModule(req); got != want { + t.Errorf("dependencyModule(%q) = %q, want %q", rawURL, got, want) + } + } +} + +func TestClientInstrumentationDoesNotMutateOriginal(t *testing.T) { + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + original := &http.Client{Transport: base, Timeout: time.Second} + instrumented := InstrumentClient(original) + if instrumented == original { + t.Fatal("InstrumentClient returned the original client") + } + if _, ok := original.Transport.(roundTripFunc); !ok { + t.Fatalf("InstrumentClient mutated the original transport to %T", original.Transport) + } + if instrumented.Timeout != original.Timeout { + t.Fatal("InstrumentClient did not preserve client settings") + } + if WrapRoundTripper(instrumented.Transport) != instrumented.Transport { + t.Fatal("WrapRoundTripper wrapped an already instrumented transport twice") + } +} diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go new file mode 100644 index 0000000000..15b9c7e50e --- /dev/null +++ b/backend/internal/pkg/xai/billing.go @@ -0,0 +1,372 @@ +package xai + +import ( + "encoding/json" + "fmt" + "math" + "net/http" + "strconv" + "strings" + "time" +) + +const ( + // CLI client identity required by cli-chat-proxy billing endpoints. + CLITokenAuthHeader = "x-xai-token-auth" + CLITokenAuthValue = "xai-grok-cli" + CLIClientVersionHeader = "x-grok-client-version" + // Keep in sync with https://x.ai/cli/stable. + CLIClientVersion = "0.2.93" + CLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + + BillingWeeklyPath = "/billing?format=credits" + BillingMonthlyPath = "/billing" + + SuperGrokLimitCents = 15_000 // $150.00 + SuperGrokHeavyLimitCents = 150_000 // $1,500.00 +) + +// BillingPeriod describes the current weekly/monthly window. +type BillingPeriod struct { + Type string `json:"type,omitempty"` + Start string `json:"start,omitempty"` + End string `json:"end,omitempty"` +} + +// BillingProductUsage is per-product usage inside the weekly credits window. +type BillingProductUsage struct { + Product string `json:"product,omitempty"` + UsagePercent *float64 `json:"usagePercent,omitempty"` +} + +// BillingConfig is the nested config object from /v1/billing responses. +type BillingConfig struct { + CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"` + CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"` + ProductUsage []BillingProductUsage `json:"productUsage,omitempty"` + MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"` + Used json.RawMessage `json:"used,omitempty"` + BillingPeriodStart string `json:"billingPeriodStart,omitempty"` + BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"` +} + +// BillingPayload is the top-level body from /v1/billing. +type BillingPayload struct { + Config *BillingConfig `json:"config,omitempty"` +} + +// BillingProductSummary is a normalized product usage row for UI. +type BillingProductSummary struct { + Product string `json:"product"` + UsagePercent *float64 `json:"usage_percent,omitempty"` +} + +// BillingSummary is the merged weekly + monthly billing view. +type BillingSummary struct { + PeriodType string `json:"period_type,omitempty"` // weekly | monthly | unknown + UsagePercent *float64 `json:"usage_percent,omitempty"` + PeriodStart string `json:"period_start,omitempty"` + PeriodEnd string `json:"period_end,omitempty"` + ProductUsage []BillingProductSummary `json:"product_usage,omitempty"` + MonthlyLimitCents *float64 `json:"monthly_limit_cents,omitempty"` + UsedCents *float64 `json:"used_cents,omitempty"` + IncludedUsedCents *float64 `json:"included_used_cents,omitempty"` + BillingPeriodStart string `json:"billing_period_start,omitempty"` + BillingPeriodEnd string `json:"billing_period_end,omitempty"` + UsedPercent *float64 `json:"used_percent,omitempty"` + Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | "" + StatusCode int `json:"status_code,omitempty"` + Source string `json:"source,omitempty"` + FetchedAt string `json:"fetched_at,omitempty"` + UpdatedAt string `json:"updated_at,omitempty"` + WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"` + MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"` + Partial bool `json:"partial,omitempty"` + FailedWindows []string `json:"failed_windows,omitempty"` +} + +// BuildBillingURL builds weekly or monthly billing URL against the CLI chat proxy. +func BuildBillingURL(formatCredits bool) string { + base := strings.TrimRight(DefaultCLIBaseURL, "/") + if formatCredits { + return base + BillingWeeklyPath + } + return base + BillingMonthlyPath +} + +// ApplyCLIBillingHeaders sets Authorization + CLI identity headers for billing GETs. +func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { + if req == nil { + return + } + token := strings.TrimSpace(accessToken) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + req.Header.Set("Accept", "application/json") + req.Header.Set("Content-Type", "application/json") + req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) + req.Header.Set(CLIClientVersionHeader, CLIClientVersion) + req.Header.Set("User-Agent", CLIUserAgent) +} + +// ParseBillingPayload unmarshals a billing API response body. +func ParseBillingPayload(body []byte) (*BillingPayload, error) { + if len(body) == 0 { + return nil, fmt.Errorf("empty billing body") + } + var payload BillingPayload + if err := json.Unmarshal(body, &payload); err != nil { + return nil, err + } + return &payload, nil +} + +// BuildBillingSummary normalizes a billing config into a UI-friendly summary. +func BuildBillingSummary(config *BillingConfig) *BillingSummary { + if config == nil { + return nil + } + summary := &BillingSummary{} + period := config.CurrentPeriod + periodType := resolvePeriodType(period) + creditUsage := cloneFloat(config.CreditUsagePercent) + + periodStart := "" + periodEnd := "" + if period != nil { + periodStart = strings.TrimSpace(period.Start) + periodEnd = strings.TrimSpace(period.End) + } + if periodStart == "" { + periodStart = strings.TrimSpace(config.BillingPeriodStart) + } + if periodEnd == "" { + periodEnd = strings.TrimSpace(config.BillingPeriodEnd) + } + + products := make([]BillingProductSummary, 0, len(config.ProductUsage)) + for _, item := range config.ProductUsage { + product := strings.TrimSpace(item.Product) + if product == "" { + continue + } + products = append(products, BillingProductSummary{ + Product: product, + UsagePercent: cloneFloat(item.UsagePercent), + }) + } + + monthlyLimit := parseCentValue(config.MonthlyLimit) + used := parseCentValue(config.Used) + billingStart := strings.TrimSpace(config.BillingPeriodStart) + billingEnd := strings.TrimSpace(config.BillingPeriodEnd) + + var includedUsed *float64 + if used != nil { + if monthlyLimit != nil && *monthlyLimit > 0 { + v := math.Min(*used, *monthlyLimit) + includedUsed = &v + } else { + includedUsed = cloneFloat(used) + } + } + + var usedPercent *float64 + if monthlyLimit != nil && *monthlyLimit > 0 && includedUsed != nil { + v := (*includedUsed / *monthlyLimit) * 100 + usedPercent = &v + } + + hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0 + hasMonthly := monthlyLimit != nil || used != nil || (!hasWeekly && billingEnd != "") + if !hasWeekly && !hasMonthly { + return nil + } + + if hasWeekly { + if periodType == "unknown" { + periodType = "weekly" + } + summary.PeriodType = periodType + summary.UsagePercent = creditUsage + summary.PeriodStart = periodStart + summary.PeriodEnd = periodEnd + } else { + // Monthly-only: do not put monthly % into UsagePercent (weekly bar field). + // Frontend weekly bar only renders when PeriodType == weekly. + summary.PeriodType = "monthly" + summary.PeriodStart = billingStart + summary.PeriodEnd = billingEnd + } + summary.ProductUsage = products + summary.MonthlyLimitCents = monthlyLimit + summary.UsedCents = used + summary.IncludedUsedCents = includedUsed + if hasMonthly { + summary.BillingPeriodStart = billingStart + summary.BillingPeriodEnd = billingEnd + } + summary.UsedPercent = usedPercent + summary.Plan = resolvePlan(monthlyLimit) + return summary +} + +// MergeBillingProbeResult updates successful billing domains while retaining +// the previous value for any domain that could not be refreshed. +func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK, monthlyOK bool) *BillingSummary { + var out BillingSummary + if previous != nil { + out = *previous + previousUpdatedAt := previous.UpdatedAt + if previousUpdatedAt == "" { + previousUpdatedAt = previous.FetchedAt + } + if out.WeeklyUpdatedAt == "" && (out.UsagePercent != nil || len(out.ProductUsage) > 0) { + out.WeeklyUpdatedAt = previousUpdatedAt + } + if out.MonthlyUpdatedAt == "" && (out.MonthlyLimitCents != nil || out.UsedPercent != nil) { + out.MonthlyUpdatedAt = previousUpdatedAt + } + } + now := time.Now().UTC().Format(time.RFC3339) + + if weeklyOK && weekly != nil { + out.PeriodType = weekly.PeriodType + out.UsagePercent = weekly.UsagePercent + out.PeriodStart = weekly.PeriodStart + out.PeriodEnd = weekly.PeriodEnd + out.ProductUsage = weekly.ProductUsage + out.WeeklyUpdatedAt = now + } + if monthlyOK && monthly != nil { + if out.PeriodType == "" { + out.PeriodType = "monthly" + } + out.MonthlyLimitCents = monthly.MonthlyLimitCents + out.UsedCents = monthly.UsedCents + out.IncludedUsedCents = monthly.IncludedUsedCents + out.BillingPeriodStart = monthly.BillingPeriodStart + out.BillingPeriodEnd = monthly.BillingPeriodEnd + out.UsedPercent = monthly.UsedPercent + out.Plan = monthly.Plan + out.MonthlyUpdatedAt = now + } + + out.Partial = !weeklyOK || !monthlyOK + out.FailedWindows = nil + if !weeklyOK { + out.FailedWindows = append(out.FailedWindows, "weekly") + } + if !monthlyOK { + out.FailedWindows = append(out.FailedWindows, "monthly") + } + if !weeklyOK && !monthlyOK && previous == nil { + return nil + } + return &out +} + +// StampBillingSummary sets fetch metadata. +func StampBillingSummary(summary *BillingSummary, statusCode int, source string) *BillingSummary { + if summary == nil { + return nil + } + now := time.Now().UTC().Format(time.RFC3339) + summary.StatusCode = statusCode + summary.Source = source + summary.FetchedAt = now + summary.UpdatedAt = now + return summary +} + +func resolvePeriodType(period *BillingPeriod) string { + if period == nil { + return "unknown" + } + raw := strings.ToLower(strings.TrimSpace(period.Type)) + if strings.Contains(raw, "weekly") { + return "weekly" + } + if strings.Contains(raw, "monthly") { + return "monthly" + } + return "unknown" +} + +func resolvePlan(monthlyLimitCents *float64) string { + if monthlyLimitCents == nil { + return "" + } + // Allow small float noise. + limit := math.Round(*monthlyLimitCents) + switch limit { + case SuperGrokLimitCents: + return "SuperGrok" + case SuperGrokHeavyLimitCents: + return "SuperGrok Heavy" + default: + return "" + } +} + +func parseCentValue(raw json.RawMessage) *float64 { + if len(raw) == 0 || string(raw) == "null" { + return nil + } + // Object form: {"val": 123} + var obj struct { + Val any `json:"val"` + } + if err := json.Unmarshal(raw, &obj); err == nil && obj.Val != nil { + return anyToFloat(obj.Val) + } + // Bare number / string + var n any + if err := json.Unmarshal(raw, &n); err != nil { + return nil + } + return anyToFloat(n) +} + +func anyToFloat(v any) *float64 { + switch n := v.(type) { + case float64: + return &n + case float32: + f := float64(n) + return &f + case int: + f := float64(n) + return &f + case int64: + f := float64(n) + return &f + case json.Number: + f, err := n.Float64() + if err != nil { + return nil + } + return &f + case string: + s := strings.TrimSpace(n) + if s == "" { + return nil + } + f, err := strconv.ParseFloat(s, 64) + if err != nil { + return nil + } + return &f + default: + return nil + } +} + +func cloneFloat(v *float64) *float64 { + if v == nil { + return nil + } + f := *v + return &f +} diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go new file mode 100644 index 0000000000..1d863f6a39 --- /dev/null +++ b/backend/internal/pkg/xai/billing_test.go @@ -0,0 +1,127 @@ +package xai + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestBuildBillingURL(t *testing.T) { + t.Parallel() + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing?format=credits", BuildBillingURL(true)) + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing", BuildBillingURL(false)) +} + +func TestApplyCLIBillingHeaders(t *testing.T) { + t.Parallel() + req, err := http.NewRequest(http.MethodGet, BuildBillingURL(true), nil) + require.NoError(t, err) + + ApplyCLIBillingHeaders(req, " token ") + + require.Equal(t, "Bearer token", req.Header.Get("Authorization")) + require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader)) + require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader)) + require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent()) +} + +func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { + t.Parallel() + + weeklyBody := []byte(`{ + "config": { + "currentPeriod": {"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}, + "creditUsagePercent": 2.0, + "productUsage": [{"product":"Api","usagePercent":2.0}] + } + }`) + monthlyBody := []byte(`{ + "config": { + "monthlyLimit": {"val": 15000}, + "used": {"val": 78}, + "billingPeriodStart": "2026-07-01T00:00:00Z", + "billingPeriodEnd": "2026-08-01T00:00:00Z" + } + }`) + + weeklyPayload, err := ParseBillingPayload(weeklyBody) + require.NoError(t, err) + monthlyPayload, err := ParseBillingPayload(monthlyBody) + require.NoError(t, err) + + weekly := BuildBillingSummary(weeklyPayload.Config) + monthly := BuildBillingSummary(monthlyPayload.Config) + require.NotNil(t, weekly) + require.NotNil(t, monthly) + require.Equal(t, "weekly", weekly.PeriodType) + require.InDelta(t, 2.0, *weekly.UsagePercent, 1e-9) + require.Equal(t, "Api", weekly.ProductUsage[0].Product) + require.Equal(t, "SuperGrok", monthly.Plan) + require.InDelta(t, 15000, *monthly.MonthlyLimitCents, 1e-9) + require.InDelta(t, 78, *monthly.UsedCents, 1e-9) + require.InDelta(t, 0.52, *monthly.UsedPercent, 1e-2) + + merged := MergeBillingProbeResult(nil, weekly, monthly, true, true) + require.Equal(t, "weekly", merged.PeriodType) + require.InDelta(t, 2.0, *merged.UsagePercent, 1e-9) + require.Equal(t, "SuperGrok", merged.Plan) + require.InDelta(t, 15000, *merged.MonthlyLimitCents, 1e-9) + require.Equal(t, "2026-08-01T00:00:00Z", merged.BillingPeriodEnd) +} + +func TestParseCentValueBareNumber(t *testing.T) { + t.Parallel() + raw, _ := json.Marshal(15000) + v := parseCentValue(raw) + require.NotNil(t, v) + require.InDelta(t, 15000, *v, 1e-9) +} + +func TestBuildBillingSummaryMonthlyOnlyKeepsWeeklyUsageEmpty(t *testing.T) { + t.Parallel() + payload, err := ParseBillingPayload([]byte(`{"config":{"monthlyLimit":{"val":15000},"used":{"val":7500},"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`)) + require.NoError(t, err) + + summary := BuildBillingSummary(payload.Config) + require.NotNil(t, summary) + require.Equal(t, "monthly", summary.PeriodType) + require.Nil(t, summary.UsagePercent) + require.InDelta(t, 50, *summary.UsedPercent, 1e-9) +} + +func TestMergeBillingProbeResultRetainsFailedWindow(t *testing.T) { + t.Parallel() + previous := &BillingSummary{ + PeriodType: "weekly", + UsagePercent: floatPointer(100), + PeriodEnd: "2026-07-16T00:00:00Z", + MonthlyLimitCents: floatPointer(15000), + UsedPercent: floatPointer(20), + BillingPeriodEnd: "2026-08-01T00:00:00Z", + WeeklyUpdatedAt: "2026-07-10T00:00:00Z", + MonthlyUpdatedAt: "2026-07-10T00:00:00Z", + FailedWindows: []string{"monthly"}, + } + monthly := &BillingSummary{ + PeriodType: "monthly", + MonthlyLimitCents: floatPointer(15000), + UsedPercent: floatPointer(30), + BillingPeriodEnd: "2026-08-01T00:00:00Z", + } + + merged := MergeBillingProbeResult(previous, nil, monthly, false, true) + require.Equal(t, "weekly", merged.PeriodType) + require.InDelta(t, 100, *merged.UsagePercent, 1e-9) + require.Equal(t, previous.WeeklyUpdatedAt, merged.WeeklyUpdatedAt) + require.InDelta(t, 30, *merged.UsedPercent, 1e-9) + require.NotEqual(t, previous.MonthlyUpdatedAt, merged.MonthlyUpdatedAt) + require.True(t, merged.Partial) + require.Equal(t, []string{"weekly"}, merged.FailedWindows) + require.Equal(t, []string{"monthly"}, previous.FailedWindows) +} + +func floatPointer(value float64) *float64 { + return &value +} diff --git a/backend/internal/pkg/xai/oauth.go b/backend/internal/pkg/xai/oauth.go index 1b3aadc08b..a8a549c9dc 100644 --- a/backend/internal/pkg/xai/oauth.go +++ b/backend/internal/pkg/xai/oauth.go @@ -191,7 +191,7 @@ func RuntimeSanity() RuntimeSanityReport { UnsafeURLOverrides: AllowUnsafeURLOverrides(), UnsafeHighConcurrency: AllowUnsafeHighConcurrency(), PublicGatewayScope: "responses_only", - ProxyPolicy: "account_proxy_optional; upstream URL allowlists enforced unless unsafe overrides are enabled", + ProxyPolicy: "account_proxy_optional; OAuth URLs use trusted-host allowlists; API-key base URLs require public HTTPS unless unsafe overrides are enabled", } } @@ -252,6 +252,19 @@ func ValidateOAuthEndpointURL(raw string) (string, error) { } func ValidateBaseURL(raw string) (string, error) { + if AllowUnsafeURLOverrides() { + return urlvalidator.ValidateURLFormat(raw, true) + } + normalized, err := urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{ + AllowPrivate: false, + }) + if err != nil { + return "", err + } + return normalizeKnownBaseURLPath(normalized) +} + +func ValidateTrustedBaseURL(raw string) (string, error) { if AllowUnsafeURLOverrides() { return urlvalidator.ValidateURLFormat(raw, true) } @@ -461,6 +474,22 @@ func BuildVideosGenerationsURL(baseURL string) (string, error) { return validatedBaseURL + "/videos/generations", nil } +func BuildVideosEditsURL(baseURL string) (string, error) { + validatedBaseURL, err := ValidatedBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base url: %w", err) + } + return validatedBaseURL + "/videos/edits", nil +} + +func BuildVideosExtensionsURL(baseURL string) (string, error) { + validatedBaseURL, err := ValidatedBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base url: %w", err) + } + return validatedBaseURL + "/videos/extensions", nil +} + func BuildVideoURL(baseURL, requestID string) (string, error) { validatedBaseURL, err := ValidatedBaseURL(baseURL) if err != nil { diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index d3d3d5cb29..39200ffc39 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -129,6 +129,14 @@ func TestBuildGrokMediaURLs(t *testing.T) { require.NoError(t, err) require.Equal(t, DefaultBaseURL+"/videos/generations", videosURL) + videoEditsURL, err := BuildVideosEditsURL(DefaultBaseURL) + require.NoError(t, err) + require.Equal(t, DefaultBaseURL+"/videos/edits", videoEditsURL) + + videoExtensionsURL, err := BuildVideosExtensionsURL(DefaultBaseURL) + require.NoError(t, err) + require.Equal(t, DefaultBaseURL+"/videos/extensions", videoExtensionsURL) + videoURL, err := BuildVideoURL(DefaultBaseURL, "req 123") require.NoError(t, err) require.Equal(t, DefaultBaseURL+"/videos/req%20123", videoURL) @@ -137,13 +145,10 @@ func TestBuildGrokMediaURLs(t *testing.T) { require.Error(t, err) } -func TestValidateXAIURLsRejectArbitraryHostsByDefault(t *testing.T) { +func TestValidateXAIURLsRejectUntrustedOAuthAndUnsafeBaseURLsByDefault(t *testing.T) { _, err := ValidateOAuthEndpointURL("https://auth.example.test/oauth2/token") require.Error(t, err) - _, err = ValidateBaseURL("https://xai.test/v1") - require.Error(t, err) - _, err = ValidateBaseURL("http://127.0.0.1:8080/v1") require.Error(t, err) @@ -151,6 +156,15 @@ func TestValidateXAIURLsRejectArbitraryHostsByDefault(t *testing.T) { require.Error(t, err) } +func TestValidateBaseURLAllowsPublicThirdPartyGrokAPI(t *testing.T) { + baseURL, err := ValidateBaseURL("https://grok.example.test/v1/") + require.NoError(t, err) + require.Equal(t, "https://grok.example.test/v1", baseURL) + + _, err = ValidateTrustedBaseURL("https://grok.example.test/v1") + require.Error(t, err) +} + func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) { t.Setenv(EnvAllowUnsafeURLOverrides, "true") @@ -182,6 +196,7 @@ func TestRuntimeSanityReportsSafeDefaults(t *testing.T) { require.False(t, report.UnsafeHighConcurrency) require.Equal(t, "responses_only", report.PublicGatewayScope) require.Contains(t, report.ProxyPolicy, "account_proxy_optional") + require.Contains(t, report.ProxyPolicy, "API-key base URLs require public HTTPS") } func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) { diff --git a/backend/internal/pkg/xai/sso_device.go b/backend/internal/pkg/xai/sso_device.go new file mode 100644 index 0000000000..e533e394d2 --- /dev/null +++ b/backend/internal/pkg/xai/sso_device.go @@ -0,0 +1,418 @@ +package xai + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "sort" + "strconv" + "strings" + "time" +) + +const ( + SSOBuildScope = "openid profile email offline_access grok-cli:access api:access conversations:read conversations:write" + SSOAccountsURL = "https://accounts.x.ai/" + SSODeviceURL = OAuthIssuer + "/oauth2/device/code" + SSOVerifyURL = OAuthIssuer + "/oauth2/device/verify" + SSOApproveURL = OAuthIssuer + "/oauth2/device/approve" + SSOTokenURL = OAuthIssuer + "/oauth2/token" + SSOConversionTimeout = 90 * time.Second + + ssoMaxAuthBody = 2 << 20 + ssoDefaultUA = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36" + ssoDefaultTokenTTL = 6 * time.Hour +) + +var ( + ErrSSOUnauthorized = errors.New("xai sso unauthorized") + ErrSSOAuthorizationDenied = errors.New("xai device authorization denied") +) + +type SSOHTTPError struct{ Status int } + +func (e SSOHTTPError) Error() string { return fmt.Sprintf("xAI OAuth HTTP %d", e.Status) } + +type SSODeviceHTTPClient interface { + Do(*http.Request) (*http.Response, error) +} + +type SSODeviceOptions struct { + HTTPClient SSODeviceHTTPClient + UserAgent string + Sleep func(context.Context, time.Duration) error +} + +type ssoDeviceFlow struct { + client SSODeviceHTTPClient + userAgent string + cookies map[string]string + sleep func(context.Context, time.Duration) error +} + +func ConvertSSOToBuild(ctx context.Context, ssoToken string, opts *SSODeviceOptions) (*TokenResponse, error) { + ssoToken = NormalizeSSOToken(ssoToken) + if ssoToken == "" { + return nil, ErrSSOUnauthorized + } + if opts == nil { + opts = &SSODeviceOptions{} + } + client := opts.HTTPClient + if client == nil { + client = &http.Client{ + Timeout: SSOConversionTimeout, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + } + } + userAgent := strings.TrimSpace(opts.UserAgent) + if userAgent == "" { + userAgent = ssoDefaultUA + } + sleep := opts.Sleep + if sleep == nil { + sleep = sleepContext + } + + flow := &ssoDeviceFlow{ + client: client, + userAgent: userAgent, + cookies: map[string]string{"sso": ssoToken, "sso-rw": ssoToken}, + sleep: sleep, + } + return flow.convert(ctx) +} + +func (f *ssoDeviceFlow) convert(ctx context.Context) (*TokenResponse, error) { + status, finalURL, _, err := f.do(ctx, http.MethodGet, SSOAccountsURL, nil) + if err != nil { + return nil, err + } + if status == http.StatusUnauthorized || strings.Contains(finalURL, "sign-in") || strings.Contains(finalURL, "sign-up") { + return nil, ErrSSOUnauthorized + } + if status < 200 || status >= 400 { + return nil, fmt.Errorf("validate Grok Web SSO: %w", SSOHTTPError{Status: status}) + } + + status, _, body, err := f.do(ctx, http.MethodPost, SSODeviceURL, url.Values{ + "client_id": {DefaultClientID}, + "scope": {SSOBuildScope}, + }) + if err != nil { + return nil, err + } + if status < 200 || status >= 300 { + return nil, fmt.Errorf("start xAI device flow: %w", SSOHTTPError{Status: status}) + } + var device struct { + DeviceCode string `json:"device_code"` + UserCode string `json:"user_code"` + VerificationURIComplete string `json:"verification_uri_complete"` + Interval int `json:"interval"` + ExpiresIn int `json:"expires_in"` + } + if err := json.Unmarshal(body, &device); err != nil { + return nil, fmt.Errorf("parse xAI device flow response: %w", err) + } + if device.DeviceCode == "" || device.UserCode == "" || !safeXAIAuthURL(device.VerificationURIComplete) { + return nil, errors.New("xAI device flow response is incomplete") + } + if device.Interval <= 0 { + device.Interval = 5 + } + if device.ExpiresIn <= 0 { + device.ExpiresIn = 1800 + } + + status, _, _, err = f.do(ctx, http.MethodGet, device.VerificationURIComplete, nil) + if err != nil { + return nil, err + } + if status < 200 || status >= 400 { + return nil, fmt.Errorf("open xAI device verification page: %w", SSOHTTPError{Status: status}) + } + + status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOVerifyURL, url.Values{"user_code": {device.UserCode}}) + if err != nil { + return nil, err + } + if status < 200 || status >= 400 { + return nil, fmt.Errorf("verify xAI device code: %w", SSOHTTPError{Status: status}) + } + if !strings.Contains(finalURL, "consent") { + return nil, errors.New("xAI device verification did not reach consent page") + } + + status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOApproveURL, url.Values{ + "user_code": {device.UserCode}, + "action": {"allow"}, + "principal_type": {"User"}, + "principal_id": {""}, + }) + if err != nil { + return nil, err + } + if status < 200 || status >= 400 { + return nil, fmt.Errorf("approve xAI device code: %w", SSOHTTPError{Status: status}) + } + if !strings.Contains(finalURL, "done") { + return nil, errors.New("xAI device approval did not reach done page") + } + + return f.pollToken(ctx, device.DeviceCode, time.Duration(device.Interval)*time.Second, time.Duration(device.ExpiresIn)*time.Second) +} + +func (f *ssoDeviceFlow) pollToken(ctx context.Context, deviceCode string, interval, expiresIn time.Duration) (*TokenResponse, error) { + if interval < time.Second { + interval = time.Second + } + deadline := time.Now().Add(minDuration(expiresIn, 75*time.Second)) + for time.Now().Before(deadline) { + if err := f.sleep(ctx, interval); err != nil { + return nil, err + } + status, _, body, err := f.do(ctx, http.MethodPost, SSOTokenURL, url.Values{ + "grant_type": {"urn:ietf:params:oauth:grant-type:device_code"}, + "client_id": {DefaultClientID}, + "device_code": {deviceCode}, + }) + if err != nil { + return nil, err + } + var payload struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + IDToken string `json:"id_token"` + TokenType string `json:"token_type"` + ExpiresIn int64 `json:"expires_in"` + Scope string `json:"scope"` + Error string `json:"error"` + ErrorDescription string `json:"error_description"` + } + if err := json.Unmarshal(body, &payload); err != nil { + return nil, fmt.Errorf("parse xAI token response: %w", err) + } + if status >= 200 && status < 300 && payload.AccessToken != "" { + if payload.ExpiresIn <= 0 { + payload.ExpiresIn = int64(ssoDefaultTokenTTL.Seconds()) + } + if payload.TokenType == "" { + payload.TokenType = "Bearer" + } + return &TokenResponse{ + AccessToken: payload.AccessToken, + RefreshToken: payload.RefreshToken, + IDToken: payload.IDToken, + TokenType: payload.TokenType, + ExpiresIn: payload.ExpiresIn, + Scope: payload.Scope, + }, nil + } + switch payload.Error { + case "authorization_pending": + continue + case "slow_down": + interval += 5 * time.Second + continue + case "access_denied", "expired_token": + return nil, ErrSSOAuthorizationDenied + default: + if status >= 400 { + return nil, fmt.Errorf("xAI token polling failed (%s): %w", firstNonEmpty(payload.ErrorDescription, payload.Error), SSOHTTPError{Status: status}) + } + return nil, fmt.Errorf("xAI token polling failed: %s", firstNonEmpty(payload.ErrorDescription, payload.Error, strconv.Itoa(status))) + } + } + return nil, errors.New("xAI device flow token polling timed out") +} + +func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form url.Values) (int, string, []byte, error) { + if !safeXAIAuthURL(endpoint) { + return 0, "", nil, errors.New("xAI OAuth URL is not trusted") + } + currentURL := endpoint + currentMethod := method + currentForm := form + for redirects := 0; redirects <= 8; redirects++ { + var body io.Reader + if currentForm != nil { + body = strings.NewReader(currentForm.Encode()) + } + request, err := http.NewRequestWithContext(ctx, currentMethod, currentURL, body) + if err != nil { + return 0, currentURL, nil, err + } + request.Header.Set("Accept", "application/json, text/html;q=0.9, */*;q=0.8") + request.Header.Set("Accept-Language", "zh-CN,zh;q=0.9,en;q=0.8") + request.Header.Set("User-Agent", f.userAgent) + if cookie := f.cookieHeader(); cookie != "" { + request.Header.Set("Cookie", cookie) + } + if currentForm != nil { + request.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + + response, err := f.client.Do(request) + if err != nil { + return 0, currentURL, nil, err + } + f.captureCookies(response) + data, readErr := io.ReadAll(io.LimitReader(response.Body, ssoMaxAuthBody+1)) + _ = response.Body.Close() + if readErr != nil { + return response.StatusCode, currentURL, nil, readErr + } + if len(data) > ssoMaxAuthBody { + return response.StatusCode, currentURL, nil, errors.New("xAI OAuth response exceeds 2 MiB") + } + if response.StatusCode < 300 || response.StatusCode > 399 { + return response.StatusCode, currentURL, data, nil + } + + location := strings.TrimSpace(response.Header.Get("Location")) + if location == "" { + return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirect missing Location") + } + base, _ := url.Parse(currentURL) + next, err := url.Parse(location) + if err != nil { + return response.StatusCode, currentURL, data, err + } + currentURL = base.ResolveReference(next).String() + if !safeXAIAuthURL(currentURL) { + return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirected to untrusted host") + } + if response.StatusCode == http.StatusSeeOther || ((response.StatusCode == http.StatusMovedPermanently || response.StatusCode == http.StatusFound) && currentMethod != http.MethodGet && currentMethod != http.MethodHead) { + currentMethod = http.MethodGet + currentForm = nil + } + } + return 0, currentURL, nil, errors.New("xAI OAuth redirected too many times") +} + +func (f *ssoDeviceFlow) captureCookies(response *http.Response) { + for _, cookie := range response.Cookies() { + name := strings.TrimSpace(cookie.Name) + value := strings.TrimSpace(cookie.Value) + if name == "" || len(name) > 128 || len(value) > 16384 || strings.ContainsAny(name+value, "\r\n\x00") { + continue + } + if cookie.MaxAge < 0 { + delete(f.cookies, name) + continue + } + f.cookies[name] = value + } +} + +func (f *ssoDeviceFlow) cookieHeader() string { + keys := make([]string, 0, len(f.cookies)) + for key := range f.cookies { + keys = append(keys, key) + } + sort.Strings(keys) + parts := make([]string, 0, len(keys)) + for _, key := range keys { + parts = append(parts, key+"="+f.cookies[key]) + } + return strings.Join(parts, "; ") +} + +func safeXAIAuthURL(raw string) bool { + parsed, err := url.Parse(raw) + if err != nil || parsed.User != nil || parsed.Hostname() == "" { + return false + } + if AllowUnsafeURLOverrides() { + return parsed.Scheme != "" && parsed.Host != "" + } + if parsed.Scheme != "https" { + return false + } + host := strings.ToLower(parsed.Hostname()) + return host == "x.ai" || strings.HasSuffix(host, ".x.ai") +} + +func NormalizeSSOToken(value string) string { + value = strings.TrimSpace(value) + if strings.HasPrefix(strings.ToLower(value), "cookie:") { + value = strings.TrimSpace(value[len("cookie:"):]) + } + for _, part := range strings.Split(value, ";") { + name, token, found := strings.Cut(strings.TrimSpace(part), "=") + if !found { + continue + } + switch strings.ToLower(strings.TrimSpace(name)) { + case "sso", "sso-rw": + return sanitizeSSOToken(token) + } + } + if token, _, found := strings.Cut(value, ";"); found { + value = strings.TrimSpace(token) + } + return sanitizeSSOToken(value) +} + +func sanitizeSSOToken(value string) string { + return strings.NewReplacer("\r", "", "\n", "", "\x00", "").Replace(strings.TrimSpace(value)) +} + +func DecodeJWTClaims(token string) map[string]any { + parts := strings.Split(token, ".") + if len(parts) < 2 { + return nil + } + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return nil + } + var claims map[string]any + if err := json.Unmarshal(payload, &claims); err != nil { + return nil + } + return claims +} + +func JWTClaimString(claims map[string]any, key string) string { + value, _ := claims[key].(string) + return strings.TrimSpace(value) +} + +func sleepContext(ctx context.Context, d time.Duration) error { + timer := time.NewTimer(d) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +func minDuration(a, b time.Duration) time.Duration { + if a <= 0 { + return b + } + if a < b { + return a + } + return b +} + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + return "" +} diff --git a/backend/internal/pkg/xai/sso_device_test.go b/backend/internal/pkg/xai/sso_device_test.go new file mode 100644 index 0000000000..27dff15f4d --- /dev/null +++ b/backend/internal/pkg/xai/sso_device_test.go @@ -0,0 +1,115 @@ +//go:build unit + +package xai + +import ( + "context" + "io" + "net/http" + "net/url" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type ssoDeviceFakeClient struct { + t *testing.T + tokenCalls int + cookieHeaders []string +} + +func (c *ssoDeviceFakeClient) Do(req *http.Request) (*http.Response, error) { + c.cookieHeaders = append(c.cookieHeaders, req.Header.Get("Cookie")) + switch req.URL.String() { + case SSOAccountsURL: + require.Equal(c.t, http.MethodGet, req.Method) + return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"session=web-session; Path=/"}}, `{}`), nil + case SSODeviceURL: + require.Equal(c.t, http.MethodPost, req.Method) + values := readSSODeviceForm(c.t, req) + require.Equal(c.t, DefaultClientID, values.Get("client_id")) + require.Equal(c.t, SSOBuildScope, values.Get("scope")) + return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"csrf=csrf-token; Path=/"}}, `{"device_code":"device-1","user_code":"USER-1","verification_uri_complete":"https://auth.x.ai/oauth2/device/complete","interval":1,"expires_in":60}`), nil + case "https://auth.x.ai/oauth2/device/complete": + require.Equal(c.t, http.MethodGet, req.Method) + return ssoDeviceResponse(http.StatusOK, nil, `ok`), nil + case SSOVerifyURL: + require.Equal(c.t, http.MethodPost, req.Method) + values := readSSODeviceForm(c.t, req) + require.Equal(c.t, "USER-1", values.Get("user_code")) + return ssoDeviceResponse(http.StatusFound, http.Header{"Location": {"/oauth2/device/consent"}}, ``), nil + case "https://auth.x.ai/oauth2/device/consent": + require.Equal(c.t, http.MethodGet, req.Method) + return ssoDeviceResponse(http.StatusOK, nil, `consent`), nil + case SSOApproveURL: + require.Equal(c.t, http.MethodPost, req.Method) + values := readSSODeviceForm(c.t, req) + require.Equal(c.t, "USER-1", values.Get("user_code")) + require.Equal(c.t, "allow", values.Get("action")) + require.Equal(c.t, "User", values.Get("principal_type")) + return ssoDeviceResponse(http.StatusSeeOther, http.Header{"Location": {"/oauth2/device/done"}}, ``), nil + case "https://auth.x.ai/oauth2/device/done": + require.Equal(c.t, http.MethodGet, req.Method) + return ssoDeviceResponse(http.StatusOK, nil, `done`), nil + case SSOTokenURL: + require.Equal(c.t, http.MethodPost, req.Method) + c.tokenCalls++ + values := readSSODeviceForm(c.t, req) + require.Equal(c.t, "urn:ietf:params:oauth:grant-type:device_code", values.Get("grant_type")) + require.Equal(c.t, "device-1", values.Get("device_code")) + return ssoDeviceResponse(http.StatusOK, nil, `{"access_token":"access-token","refresh_token":"refresh-token","id_token":"id-token","token_type":"Bearer","expires_in":3600,"scope":"`+SSOBuildScope+`"}`), nil + default: + c.t.Fatalf("unexpected request: %s %s", req.Method, req.URL.String()) + return nil, nil + } +} + +func TestConvertSSOToBuildCompletesDeviceFlow(t *testing.T) { + t.Setenv(EnvClientID, "") + client := &ssoDeviceFakeClient{t: t} + token, err := ConvertSSOToBuild(context.Background(), "sso=sso-token; ignored=1", &SSODeviceOptions{ + HTTPClient: client, + Sleep: func(context.Context, time.Duration) error { + return nil + }, + }) + + require.NoError(t, err) + require.Equal(t, "access-token", token.AccessToken) + require.Equal(t, "refresh-token", token.RefreshToken) + require.Equal(t, "id-token", token.IDToken) + require.Equal(t, SSOBuildScope, token.Scope) + require.Equal(t, 1, client.tokenCalls) + require.Contains(t, client.cookieHeaders[0], "sso=sso-token") + require.Contains(t, client.cookieHeaders[0], "sso-rw=sso-token") + require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "session=web-session") + require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "csrf=csrf-token") +} + +func TestNormalizeSSOTokenAcceptsCookieHeader(t *testing.T) { + require.Equal(t, "token-1", NormalizeSSOToken("Cookie: foo=bar; sso=token-1; sso-rw=token-2")) + require.Equal(t, "token-2", NormalizeSSOToken("sso-rw=token-2; foo=bar")) + require.Equal(t, "raw-token", NormalizeSSOToken(" raw-token ; ignored=1")) +} + +func ssoDeviceResponse(status int, header http.Header, body string) *http.Response { + if header == nil { + header = http.Header{} + } + return &http.Response{ + StatusCode: status, + Header: header, + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func readSSODeviceForm(t *testing.T, req *http.Request) url.Values { + t.Helper() + data, err := io.ReadAll(req.Body) + require.NoError(t, err) + values, err := url.ParseQuery(string(data)) + require.NoError(t, err) + return values +} diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index c3c2f6708a..8eb819aeab 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -61,6 +61,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{ var schedulerNeutralExtraKeys = map[string]struct{}{ "codex_usage_updated_at": {}, + "grok_billing_snapshot": {}, "session_window_utilization": {}, } @@ -1286,6 +1287,38 @@ func (r *accountRepository) SetRateLimited(ctx context.Context, id int64, resetA return nil } +// SetRateLimitedIfLater atomically extends an account-level rate limit. Grok +// requests may finish concurrently, so an older response must not overwrite a +// later reset boundary observed by another request or instance. +func (r *accountRepository) SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error { + now := time.Now() + updated, err := r.client.Account.Update(). + Where( + dbaccount.IDEQ(id), + dbaccount.Or( + dbaccount.RateLimitResetAtIsNil(), + dbaccount.RateLimitResetAtLT(resetAt), + ), + ). + SetRateLimitedAt(now). + SetRateLimitResetAt(resetAt). + Save(ctx) + if err != nil { + return err + } + if updated == 0 { + // This instance may not have observed the later value written elsewhere. + // Refresh its local scheduler snapshot even though no outbox event is needed. + r.syncSchedulerAccountSnapshot(ctx, id) + return nil + } + if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { + logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue extended rate limit failed: account=%d err=%v", id, err) + } + r.syncSchedulerAccountSnapshot(ctx, id) + return nil +} + func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error { if scope == "" { return nil @@ -1521,7 +1554,7 @@ func (r *accountRepository) SetSchedulable(ctx context.Context, id int64, schedu } func (r *accountRepository) AutoPauseExpiredAccounts(ctx context.Context, now time.Time) (int64, error) { - result, err := r.sql.ExecContext(ctx, ` + rows, err := r.sql.QueryContext(ctx, ` UPDATE accounts SET schedulable = FALSE, updated_at = NOW() @@ -1530,20 +1563,35 @@ func (r *accountRepository) AutoPauseExpiredAccounts(ctx context.Context, now ti AND auto_pause_on_expired = TRUE AND expires_at IS NOT NULL AND expires_at <= $1 + RETURNING id `, now) if err != nil { return 0, err } - rows, err := result.RowsAffected() - if err != nil { + defer func() { + _ = rows.Close() + }() + + accountIDs := make([]int64, 0) + for rows.Next() { + var accountID int64 + if err := rows.Scan(&accountID); err != nil { + return 0, err + } + accountIDs = append(accountIDs, accountID) + } + if err := rows.Err(); err != nil { return 0, err } - if rows > 0 { - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventFullRebuild, nil, nil, nil); err != nil { - logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue auto pause rebuild failed: err=%v", err) + + if len(accountIDs) > 0 { + // 只刷新本次暂停的账号及其所属分组,避免少量账号到期触发所有调度桶重建。 + payload := map[string]any{"account_ids": accountIDs} + if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil { + logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue auto pause account changes failed: err=%v", err) } } - return rows, nil + return int64(len(accountIDs)), nil } func (r *accountRepository) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error { diff --git a/backend/internal/repository/account_repo_auto_pause_test.go b/backend/internal/repository/account_repo_auto_pause_test.go new file mode 100644 index 0000000000..0eb48a296d --- /dev/null +++ b/backend/internal/repository/account_repo_auto_pause_test.go @@ -0,0 +1,72 @@ +package repository + +import ( + "context" + "database/sql/driver" + "encoding/json" + "reflect" + "regexp" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +type accountIDsPayloadMatcher struct { + want []int64 +} + +func (m accountIDsPayloadMatcher) Match(value driver.Value) bool { + raw, ok := value.([]byte) + if !ok { + return false + } + var payload struct { + AccountIDs []int64 `json:"account_ids"` + } + if err := json.Unmarshal(raw, &payload); err != nil { + return false + } + return reflect.DeepEqual(m.want, payload.AccountIDs) +} + +func TestAutoPauseExpiredAccountsEnqueuesAffectedAccounts(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + + now := time.Now() + mock.ExpectQuery(`(?s)UPDATE accounts.*RETURNING id`). + WithArgs(now). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(int64(11)).AddRow(int64(29))) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload)")). + WithArgs(service.SchedulerOutboxEventAccountBulkChanged, nil, nil, accountIDsPayloadMatcher{want: []int64{11, 29}}). + WillReturnResult(sqlmock.NewResult(1, 1)) + + repo := newAccountRepositoryWithSQL(nil, db, nil) + updated, err := repo.AutoPauseExpiredAccounts(context.Background(), now) + + require.NoError(t, err) + require.EqualValues(t, 2, updated) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAutoPauseExpiredAccountsSkipsOutboxWithoutChanges(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + + now := time.Now() + mock.ExpectQuery(`(?s)UPDATE accounts.*RETURNING id`). + WithArgs(now). + WillReturnRows(sqlmock.NewRows([]string{"id"})) + + repo := newAccountRepositoryWithSQL(nil, db, nil) + updated, err := repo.AutoPauseExpiredAccounts(context.Background(), now) + + require.NoError(t, err) + require.Zero(t, updated) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/account_repo_grok_billing_test.go b/backend/internal/repository/account_repo_grok_billing_test.go new file mode 100644 index 0000000000..fb41ae5ffa --- /dev/null +++ b/backend/internal/repository/account_repo_grok_billing_test.go @@ -0,0 +1,16 @@ +package repository + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGrokBillingSnapshotIsSchedulerNeutral(t *testing.T) { + t.Parallel() + + require.True(t, isSchedulerNeutralExtraKey("grok_billing_snapshot")) + require.False(t, shouldEnqueueSchedulerOutboxForExtraUpdates(map[string]any{ + "grok_billing_snapshot": map[string]any{"usage_percent": 50}, + })) +} diff --git a/backend/internal/repository/account_repo_integration_test.go b/backend/internal/repository/account_repo_integration_test.go index de7fd3a9f3..f82ad4ca28 100644 --- a/backend/internal/repository/account_repo_integration_test.go +++ b/backend/internal/repository/account_repo_integration_test.go @@ -703,6 +703,25 @@ func (s *AccountRepoSuite) TestSetRateLimited() { s.Require().WithinDuration(resetAt, *got.RateLimitResetAt, time.Second) } +func (s *AccountRepoSuite) TestSetRateLimitedIfLaterDoesNotShortenReset() { + account := mustCreateAccount(s.T(), s.client, &service.Account{Name: "acc-rl-monotonic"}) + later := time.Now().Add(30 * time.Minute).UTC().Truncate(time.Second) + earlier := time.Now().Add(5 * time.Minute).UTC().Truncate(time.Second) + cacheRecorder := &schedulerCacheRecorder{} + s.repo.schedulerCache = cacheRecorder + + s.Require().NoError(s.repo.SetRateLimitedIfLater(s.ctx, account.ID, later)) + s.Require().NoError(s.repo.SetRateLimitedIfLater(s.ctx, account.ID, earlier)) + + got, err := s.repo.GetByID(s.ctx, account.ID) + s.Require().NoError(err) + s.Require().NotNil(got.RateLimitResetAt) + s.Require().WithinDuration(later, *got.RateLimitResetAt, time.Second) + s.Require().Len(cacheRecorder.setAccounts, 2) + s.Require().NotNil(cacheRecorder.setAccounts[1].RateLimitResetAt) + s.Require().WithinDuration(later, *cacheRecorder.setAccounts[1].RateLimitResetAt, time.Second) +} + func (s *AccountRepoSuite) TestClearRateLimit() { account := mustCreateAccount(s.T(), s.client, &service.Account{Name: "acc-clear"}) until := time.Now().Add(1 * time.Hour) diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index d348ef29b7..4c0edf72b4 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -190,6 +190,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, + group.FieldWebSearchPricePerCall, group.FieldClaudeCodeOnly, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, @@ -524,17 +525,19 @@ func (r *apiKeyRepository) latestUsageLogIPs(ctx context.Context, apiKeyIDs []in func latestUsageLogIPsQuery(apiKeyIDs []int64, dialectName string) (string, []any) { if dialectName == dialect.Postgres { + // Keep each key lookup bounded to one ordered index probe instead of ranking its full history. return ` - SELECT api_key_id, ip_address - FROM ( - SELECT api_key_id, ip_address, - ROW_NUMBER() OVER (PARTITION BY api_key_id ORDER BY created_at DESC, id DESC) AS rn - FROM usage_logs - WHERE api_key_id = ANY($1::bigint[]) - AND ip_address IS NOT NULL - AND ip_address <> '' - ) ranked - WHERE rn = 1`, []any{pq.Array(apiKeyIDs)} + SELECT requested.api_key_id, latest.ip_address + FROM unnest($1::bigint[]) AS requested(api_key_id) + CROSS JOIN LATERAL ( + SELECT ul.ip_address + FROM usage_logs AS ul + WHERE ul.api_key_id = requested.api_key_id + AND ul.ip_address IS NOT NULL + AND ul.ip_address <> '' + ORDER BY ul.created_at DESC, ul.id DESC + LIMIT 1 + ) AS latest`, []any{pq.Array(apiKeyIDs)} } placeholders := make([]string, len(apiKeyIDs)) @@ -943,6 +946,7 @@ func groupEntityToService(g *dbent.Group) *service.Group { VideoPrice480P: g.VideoPrice480p, VideoPrice720P: g.VideoPrice720p, VideoPrice1080P: g.VideoPrice1080p, + WebSearchPricePerCall: g.WebSearchPricePerCall, DefaultValidityDays: g.DefaultValidityDays, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, diff --git a/backend/internal/repository/api_key_repo_last_used_unit_test.go b/backend/internal/repository/api_key_repo_last_used_unit_test.go index 839eda7f75..dbdf653f8a 100644 --- a/backend/internal/repository/api_key_repo_last_used_unit_test.go +++ b/backend/internal/repository/api_key_repo_last_used_unit_test.go @@ -3,6 +3,7 @@ package repository import ( "context" "database/sql" + "strings" "testing" "time" @@ -125,6 +126,20 @@ func TestAPIKeyRepositoryListByUserIDAttachesLastUsedIP(t *testing.T) { require.Nil(t, byID[noLogs.ID].LastUsedIP) } +func TestLatestUsageLogIPsQueryPostgresUsesPerKeyLateralLookup(t *testing.T) { + query, args := latestUsageLogIPsQuery([]int64{11, 22}, dialect.Postgres) + normalizedQuery := strings.Join(strings.Fields(query), " ") + + require.Contains(t, normalizedQuery, "FROM unnest($1::bigint[]) AS requested(api_key_id)") + require.Contains(t, normalizedQuery, "CROSS JOIN LATERAL") + require.Contains(t, normalizedQuery, "WHERE ul.api_key_id = requested.api_key_id") + require.Contains(t, normalizedQuery, "AND ul.ip_address IS NOT NULL") + require.Contains(t, normalizedQuery, "AND ul.ip_address <> ''") + require.Contains(t, normalizedQuery, "ORDER BY ul.created_at DESC, ul.id DESC LIMIT 1") + require.NotContains(t, normalizedQuery, "ROW_NUMBER") + require.Len(t, args, 1) +} + func TestAPIKeyRepository_CreateWithLastUsedAt(t *testing.T) { repo, client := newAPIKeyRepoSQLite(t) ctx := context.Background() diff --git a/backend/internal/repository/backup_s3_store.go b/backend/internal/repository/backup_s3_store.go index 5d419f574b..2104e1e5d7 100644 --- a/backend/internal/repository/backup_s3_store.go +++ b/backend/internal/repository/backup_s3_store.go @@ -13,6 +13,7 @@ import ( "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/service" ) @@ -63,12 +64,14 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader, return 0, fmt.Errorf("read body: %w", err) } + finish := servertiming.ObserveDependency(ctx, "s3") _, err = s.client.PutObject(ctx, &s3.PutObjectInput{ Bucket: &s.bucket, Key: &key, Body: bytes.NewReader(data), ContentType: &contentType, }) + finish() if err != nil { return 0, fmt.Errorf("S3 PutObject: %w", err) } @@ -76,10 +79,12 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader, } func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser, error) { + finish := servertiming.ObserveDependency(ctx, "s3") result, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: &s.bucket, Key: &key, }) + finish() if err != nil { return nil, fmt.Errorf("S3 GetObject: %w", err) } @@ -87,10 +92,12 @@ func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser } func (s *S3BackupStore) Delete(ctx context.Context, key string) error { + finish := servertiming.ObserveDependency(ctx, "s3") _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: &s.bucket, Key: &key, }) + finish() return err } @@ -107,9 +114,11 @@ func (s *S3BackupStore) PresignURL(ctx context.Context, key string, expiry time. } func (s *S3BackupStore) HeadBucket(ctx context.Context) error { + finish := servertiming.ObserveDependency(ctx, "s3") _, err := s.client.HeadBucket(ctx, &s3.HeadBucketInput{ Bucket: &s.bucket, }) + finish() if err != nil { return fmt.Errorf("S3 HeadBucket failed: %w", err) } diff --git a/backend/internal/repository/claude_oauth_service.go b/backend/internal/repository/claude_oauth_service.go index 5c5f27c86a..ec2d426ecb 100644 --- a/backend/internal/repository/claude_oauth_service.go +++ b/backend/internal/repository/claude_oauth_service.go @@ -276,5 +276,5 @@ func createReqClient(proxyURL string) (*req.Client, error) { client.SetProxyURL(trimmed) } - return client, nil + return instrumentReqClient(client), nil } diff --git a/backend/internal/repository/concurrency_cache.go b/backend/internal/repository/concurrency_cache.go index b657c1ce8f..5341d411e9 100644 --- a/backend/internal/repository/concurrency_cache.go +++ b/backend/internal/repository/concurrency_cache.go @@ -30,6 +30,10 @@ const ( userSlotKeyPrefix = "concurrency:user:" // 格式: concurrency:api_key:{apiKeyID} apiKeySlotKeyPrefix = "concurrency:api_key:" + // API-key-scoped client WebSocket ingress leases use a shorter TTL than + // ordinary request slots, because idle ingress sessions do not hold a turn slot. + openAIWSIngressLeaseKeyPrefix = "concurrency:openai_ws_ingress:api_key:" + openAIWSIngressLeaseTTLSeconds = 60 // 等待队列计数器格式: concurrency:wait:{userID} waitQueueKeyPrefix = "concurrency:wait:" // 账号级等待队列计数器格式: wait:account:{accountID} @@ -138,6 +142,49 @@ var ( return 1 `) + // acquireOpenAIWSIngressLeaseScript atomically reaps crashed members and + // acquires or refreshes one API-key-scoped ingress lease using Redis TIME. + acquireOpenAIWSIngressLeaseScript = redis.NewScript(` + redis.replicate_commands() + local key = KEYS[1] + local maxConnections = tonumber(ARGV[1]) + local ttl = tonumber(ARGV[2]) + local leaseID = ARGV[3] + local now = tonumber(redis.call('TIME')[1]) + local expireBefore = now - ttl + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + if redis.call('ZSCORE', key, leaseID) ~= false then + redis.call('ZADD', key, now, leaseID) + redis.call('EXPIRE', key, ttl) + return 1 + end + if redis.call('ZCARD', key) < maxConnections then + redis.call('ZADD', key, now, leaseID) + redis.call('EXPIRE', key, ttl) + return 1 + end + return 0 + `) + + // refreshOpenAIWSIngressLeaseScript does not recreate a missing member: a + // process that lost its lease must terminate its local WebSocket instead of + // silently continuing beyond the distributed cap. + refreshOpenAIWSIngressLeaseScript = redis.NewScript(` + redis.replicate_commands() + local key = KEYS[1] + local ttl = tonumber(ARGV[1]) + local leaseID = ARGV[2] + local now = tonumber(redis.call('TIME')[1]) + local expireBefore = now - ttl + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + if redis.call('ZSCORE', key, leaseID) == false then + return 0 + end + redis.call('ZADD', key, now, leaseID) + 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 @@ -283,6 +330,10 @@ func apiKeySlotKey(apiKeyID int64) string { return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) } +func openAIWSIngressLeaseKey(apiKeyID int64) string { + return fmt.Sprintf("%s%d", openAIWSIngressLeaseKeyPrefix, apiKeyID) +} + func waitQueueKey(userID int64) string { return fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) } @@ -623,6 +674,48 @@ func (c *concurrencyCache) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64 return c.rdb.ZRem(ctx, key, requestID).Err() } +func (c *concurrencyCache) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) { + if c == nil || c.rdb == nil || apiKeyID <= 0 || maxConnections <= 0 || leaseID == "" { + return false, nil + } + result, err := acquireOpenAIWSIngressLeaseScript.Run( + ctx, + c.rdb, + []string{openAIWSIngressLeaseKey(apiKeyID)}, + maxConnections, + openAIWSIngressLeaseTTLSeconds, + leaseID, + ).Int() + if err != nil { + return false, err + } + return result == 1, nil +} + +func (c *concurrencyCache) RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error) { + if c == nil || c.rdb == nil || apiKeyID <= 0 || leaseID == "" { + return false, nil + } + result, err := refreshOpenAIWSIngressLeaseScript.Run( + ctx, + c.rdb, + []string{openAIWSIngressLeaseKey(apiKeyID)}, + openAIWSIngressLeaseTTLSeconds, + leaseID, + ).Int() + if err != nil { + return false, err + } + return result == 1, nil +} + +func (c *concurrencyCache) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error { + if c == nil || c.rdb == nil || apiKeyID <= 0 || leaseID == "" { + return nil + } + return c.rdb.ZRem(ctx, openAIWSIngressLeaseKey(apiKeyID), leaseID).Err() +} + func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { if len(apiKeyIDs) == 0 { return map[int64]int{}, nil diff --git a/backend/internal/repository/concurrency_cache_integration_test.go b/backend/internal/repository/concurrency_cache_integration_test.go index f7e27d1118..02821159ee 100644 --- a/backend/internal/repository/concurrency_cache_integration_test.go +++ b/backend/internal/repository/concurrency_cache_integration_test.go @@ -50,6 +50,53 @@ func (s *ConcurrencyCacheSuite) apiKeyConcurrencyCache() apiKeyConcurrencyCacheF return cache } +func (s *ConcurrencyCacheSuite) TestOpenAIWSIngressAPIKeySlot_HardLimitRefreshAndRelease() { + apiKeyID := int64(9011) + firstLeaseID := "ingress-first" + secondLeaseID := "ingress-second" + + ok, err := s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, firstLeaseID) + require.NoError(s.T(), err) + require.True(s.T(), ok) + + ok, err = s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, secondLeaseID) + require.NoError(s.T(), err) + require.False(s.T(), ok, "a second live session must not exceed the API key limit") + + ok, err = s.rawCache.RefreshOpenAIWSIngressLease(s.ctx, apiKeyID, firstLeaseID) + require.NoError(s.T(), err) + require.True(s.T(), ok, "the current owner must be able to refresh its lease") + + require.NoError(s.T(), s.rawCache.ReleaseOpenAIWSIngressLease(s.ctx, apiKeyID, firstLeaseID)) + ok, err = s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, secondLeaseID) + require.NoError(s.T(), err) + require.True(s.T(), ok, "released capacity must become available immediately") +} + +func (s *ConcurrencyCacheSuite) TestOpenAIWSIngressAPIKeySlot_ReapsCrashedLeaseWithoutDeletingLiveOtherInstance() { + apiKeyID := int64(9012) + key := openAIWSIngressLeaseKey(apiKeyID) + now, err := s.rawCache.redisUnixSeconds(s.ctx) + require.NoError(s.T(), err) + require.NoError(s.T(), s.rdb.ZAdd(s.ctx, key, + redis.Z{Score: float64(now - openAIWSIngressLeaseTTLSeconds - 1), Member: "crashed-instance"}, + redis.Z{Score: float64(now), Member: "live-other-instance"}, + ).Err()) + require.NoError(s.T(), s.rdb.Expire(s.ctx, key, time.Duration(openAIWSIngressLeaseTTLSeconds)*time.Second).Err()) + + ok, err := s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 2, "new-instance") + require.NoError(s.T(), err) + require.True(s.T(), ok, "the crashed member should be reaped before enforcing the limit") + + _, err = s.rdb.ZScore(s.ctx, key, "crashed-instance").Result() + require.ErrorIs(s.T(), err, redis.Nil) + _, err = s.rdb.ZScore(s.ctx, key, "live-other-instance").Result() + require.NoError(s.T(), err, "a live lease owned by another instance must be preserved") + count, err := s.rdb.ZCard(s.ctx, key).Result() + require.NoError(s.T(), err) + require.Equal(s.T(), int64(2), count) +} + func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() { accountID := int64(10) reqID1, reqID2, reqID3 := "req1", "req2", "req3" diff --git a/backend/internal/repository/ent.go b/backend/internal/repository/ent.go index 64d321924d..3abb528e98 100644 --- a/backend/internal/repository/ent.go +++ b/backend/internal/repository/ent.go @@ -15,7 +15,7 @@ import ( "entgo.io/ent/dialect" entsql "entgo.io/ent/dialect/sql" - _ "github.com/lib/pq" // PostgreSQL 驱动,通过副作用导入注册驱动 + "github.com/lib/pq" ) // InitEnt 初始化 Ent ORM 客户端并返回客户端实例和底层的 *sql.DB。 @@ -48,9 +48,19 @@ func InitEnt(cfg *config.Config) (*ent.Client, *sql.DB, error) { // 使用 Ent 的 SQL 驱动打开 PostgreSQL 连接。 // dialect.Postgres 指定使用 PostgreSQL 方言进行 SQL 生成。 - drv, err := entsql.Open(dialect.Postgres, dsn) - if err != nil { - return nil, nil, err + var drv *entsql.Driver + if cfg.Server.EnableServerTiming { + connector, err := pq.NewConnector(dsn) + if err != nil { + return nil, nil, err + } + drv = entsql.OpenDB(dialect.Postgres, sql.OpenDB(newServerTimingConnector(connector))) + } else { + var err error + drv, err = entsql.Open(dialect.Postgres, dsn) + if err != nil { + return nil, nil, err + } } applyDBPoolSettings(drv.DB(), cfg) diff --git a/backend/internal/repository/grok_oauth_client.go b/backend/internal/repository/grok_oauth_client.go index 435ced5a65..6c9c2c407f 100644 --- a/backend/internal/repository/grok_oauth_client.go +++ b/backend/internal/repository/grok_oauth_client.go @@ -2,12 +2,14 @@ package repository import ( "context" + "errors" "net/http" "net/url" "strings" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + sharedhttp "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/util/logredact" @@ -88,6 +90,21 @@ func (c *grokOAuthClient) RefreshToken(ctx context.Context, refreshToken, proxyU return &tokenResp, nil } +func (c *grokOAuthClient) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) { + client, err := createGrokSSOHTTPClient(proxyURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_SSO_CLIENT_INIT_FAILED", "create HTTP client: %v", err) + } + + requestCtx, cancel := context.WithTimeout(ctx, xai.SSOConversionTimeout) + defer cancel() + tokenResp, err := xai.ConvertSSOToBuild(requestCtx, ssoToken, &xai.SSODeviceOptions{HTTPClient: client}) + if err != nil { + return nil, grokSSOConversionError(err) + } + return tokenResp, nil +} + func createGrokReqClient(proxyURL string) (*req.Client, error) { return getSharedReqClient(reqClientOptions{ ProxyURL: proxyURL, @@ -95,6 +112,43 @@ func createGrokReqClient(proxyURL string) (*req.Client, error) { }) } +func createGrokSSOHTTPClient(proxyURL string) (*http.Client, error) { + client, err := sharedhttp.GetClient(sharedhttp.Options{ + ProxyURL: proxyURL, + Timeout: xai.SSOConversionTimeout, + ResponseHeaderTimeout: 30 * time.Second, + }) + if err != nil { + return nil, err + } + clone := *client + clone.CheckRedirect = func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + } + return &clone, nil +} + +func grokSSOConversionError(err error) error { + if errors.Is(err, xai.ErrSSOUnauthorized) { + return infraerrors.New(http.StatusUnauthorized, "GROK_SSO_UNAUTHORIZED", "Grok Web SSO cookie is invalid or expired") + } + if errors.Is(err, xai.ErrSSOAuthorizationDenied) { + return infraerrors.New(http.StatusForbidden, "GROK_SSO_AUTHORIZATION_DENIED", "xAI device authorization was denied or expired") + } + var statusErr xai.SSOHTTPError + if errors.As(err, &statusErr) { + statusCode := http.StatusBadGateway + if statusErr.Status == http.StatusForbidden { + statusCode = http.StatusForbidden + } + return infraerrors.Newf(statusCode, "GROK_SSO_UPSTREAM_FAILED", "xAI SSO conversion failed: %v", err) + } + if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) { + return infraerrors.Newf(http.StatusGatewayTimeout, "GROK_SSO_TIMEOUT", "xAI SSO conversion timed out: %v", err) + } + return infraerrors.Newf(http.StatusBadGateway, "GROK_SSO_CONVERSION_FAILED", "xAI SSO conversion failed: %v", err) +} + func grokOAuthStatusError(code, message string, resp *req.Response) error { statusCode := http.StatusBadGateway errorCode := code diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 37529c60be..47efb72dd8 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -63,6 +63,7 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetNillableVideoPrice480p(groupIn.VideoPrice480P). SetNillableVideoPrice720p(groupIn.VideoPrice720P). SetNillableVideoPrice1080p(groupIn.VideoPrice1080P). + SetNillableWebSearchPricePerCall(groupIn.WebSearchPricePerCall). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetNillableFallbackGroupID(groupIn.FallbackGroupID). @@ -215,6 +216,11 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er } else { builder = builder.ClearVideoPrice1080p() } + if groupIn.WebSearchPricePerCall != nil { + builder = builder.SetWebSearchPricePerCall(*groupIn.WebSearchPricePerCall) + } else { + builder = builder.ClearWebSearchPricePerCall() + } // 处理 FallbackGroupID:nil 时清除,否则设置 if groupIn.FallbackGroupID != nil { diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index eac60d6268..58b1d345d4 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -13,6 +13,7 @@ import ( "net" "net/http" "net/url" + "os" "strings" "sync" "sync/atomic" @@ -20,13 +21,16 @@ import ( "github.com/andybalholm/brotli" "github.com/klauspost/compress/zstd" + "golang.org/x/net/http2" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" + "golang.org/x/mod/semver" ) // 默认配置常量 @@ -57,6 +61,20 @@ const ( defaultOpenAIHTTP2FallbackErrorThreshold = 2 defaultOpenAIHTTP2FallbackWindow = 60 * time.Second defaultOpenAIHTTP2FallbackTTL = 10 * time.Minute + // OpenAI HTTP/2 连接健康探测:Codex 上游改走 HTTP/2 后,池化连接被代理/NAT + // 静默掐断会成为“死连接”(两端都以为存活),请求落上去会挂到 TCP 重传超时 + // (分钟级)。Go 的 http2.Transport 默认 ReadIdleTimeout=0(不发健康 PING), + // 无法检测。启用主动 PING 探测:连接空闲 ReadIdleTimeout 后发 PING,PingTimeout + // 内无响应即判定死连接并关闭,从源头避免请求挂在死连接上。 + openAIHTTP2ReadIdleTimeout = 15 * time.Second + openAIHTTP2PingTimeout = 15 * time.Second + + // The Grok CLI proxy rejects requests that do not identify a supported + // client version. Keep a known-good stable version in the binary while + // allowing operators to bump it without waiting for a Sub2API release. + grokCLIProxyHost = "cli-chat-proxy.grok.com" + grokCLIStableVersion = "0.2.93" + grokCLIVersionOverride = "XAI_GROK_CLI_VERSION" ) const ( @@ -161,6 +179,7 @@ func NewHTTPUpstream(cfg *config.Config) service.HTTPUpstream { // - 调用方必须关闭 resp.Body,否则会导致 inFlight 计数泄漏 // - inFlight > 0 的客户端不会被淘汰,确保活跃请求不被中断 func (s *httpUpstreamService) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + applyGrokCLIProxyHeaders(req) if err := s.validateRequestHost(req); err != nil { return nil, err } @@ -176,7 +195,7 @@ func (s *httpUpstreamService) Do(req *http.Request, proxyURL string, accountID i } // 执行请求 - resp, err := entry.client.Do(req) + resp, err := servertiming.Do(entry.client, req) if err != nil { s.recordOpenAIHTTP2Failure(profile, entry.protocolMode, entry.proxyKey, err) // 请求失败,立即减少计数 @@ -207,6 +226,7 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco if profile == nil { return s.Do(req, proxyURL, accountID, accountConcurrency) } + applyGrokCLIProxyHeaders(req) upstreamProfile := service.HTTPUpstreamProfileDefault if req != nil { upstreamProfile = service.HTTPUpstreamProfileFromContext(req.Context()) @@ -232,7 +252,7 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco return nil, err } - resp, err := entry.client.Do(req) + resp, err := servertiming.Do(entry.client, req) if err != nil { atomic.AddInt64(&entry.inFlight, -1) atomic.StoreInt64(&entry.lastUsed, time.Now().UnixNano()) @@ -250,6 +270,34 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco return resp, nil } +// applyGrokCLIProxyHeaders applies the official Grok Build client identity at +// the final shared transport boundary. Keying this behavior to the exact CLI +// proxy host keeps direct api.x.ai traffic unchanged and automatically covers +// Responses, Chat Completions, media, quota probes, and account tests. +func applyGrokCLIProxyHeaders(req *http.Request) { + if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), grokCLIProxyHost) { + return + } + if req.Header == nil { + req.Header = make(http.Header) + } + version := strings.TrimSpace(os.Getenv(grokCLIVersionOverride)) + if !isSupportedGrokCLIVersion(version) { + version = grokCLIStableVersion + } + req.Header.Set("X-XAI-Token-Auth", "xai-grok-cli") + req.Header.Set("x-grok-client-version", version) + req.Header.Set("User-Agent", "xai-grok-workspace/"+version) +} + +func isSupportedGrokCLIVersion(version string) bool { + canonical := "v" + version + minimum := "v" + grokCLIStableVersion + return semver.IsValid(canonical) && + semver.Canonical(canonical) == canonical && + semver.Compare(canonical, minimum) >= 0 +} + // acquireClientWithTLS 获取或创建带 TLS 指纹的客户端 func (s *httpUpstreamService) acquireClientWithTLS(proxyURL string, accountID int64, accountConcurrency int, profile *tlsfingerprint.Profile, upstreamProfile service.HTTPUpstreamProfile) (*upstreamClientEntry, error) { return s.getClientEntryWithTLS(proxyURL, accountID, accountConcurrency, profile, upstreamProfile, true, true) @@ -1062,6 +1110,11 @@ func buildUpstreamTransport(settings poolSettings, proxyURL *url.URL, protocolMo switch protocolMode { case upstreamProtocolModeOpenAIH2: transport.ForceAttemptHTTP2 = true + // 显式配置 http2 并启用 PING 健康探测,剔除代理/NAT 静默掐断的死连接, + // 避免请求挂在死连接上直到 TCP 重传超时(分钟级)。 + if _, err := enableOpenAIHTTP2KeepAlive(transport); err != nil { + return nil, err + } case upstreamProtocolModeOpenAIH1: transport.ForceAttemptHTTP2 = false transport.TLSNextProto = make(map[string]func(string, *tls.Conn) http.RoundTripper) @@ -1076,6 +1129,22 @@ func buildUpstreamTransport(settings poolSettings, proxyURL *url.URL, protocolMo return transport, nil } +// enableOpenAIHTTP2KeepAlive 在 http.Transport 上显式配置 HTTP/2 并启用连接健康探测。 +// Go 默认惰性配置 http2 且 ReadIdleTimeout=0(不发健康 PING),无法检测被代理/NAT +// 静默掐断的死连接。此处主动设置 ReadIdleTimeout/PingTimeout,让死连接被提前 PING +// 出并关闭,请求得以重建连接而非挂到 TCP 重传超时。返回底层 *http2.Transport 便于测试。 +func enableOpenAIHTTP2KeepAlive(transport *http.Transport) (*http2.Transport, error) { + h2, err := http2.ConfigureTransports(transport) + if err != nil { + return nil, err + } + if h2 != nil { + h2.ReadIdleTimeout = openAIHTTP2ReadIdleTimeout + h2.PingTimeout = openAIHTTP2PingTimeout + } + return h2, nil +} + // buildUpstreamTransportWithTLSFingerprint 构建带 TLS 指纹伪装的 Transport // 使用 utls 库模拟 Claude CLI 的 TLS 指纹 // diff --git a/backend/internal/repository/http_upstream_http2_keepalive_test.go b/backend/internal/repository/http_upstream_http2_keepalive_test.go new file mode 100644 index 0000000000..ead61da3b2 --- /dev/null +++ b/backend/internal/repository/http_upstream_http2_keepalive_test.go @@ -0,0 +1,67 @@ +package repository + +import ( + "net/http" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func http2KeepAliveTestPoolSettings() poolSettings { + return poolSettings{ + maxIdleConns: 10, + maxIdleConnsPerHost: 5, + maxConnsPerHost: 10, + idleConnTimeout: 90 * time.Second, + responseHeaderTimeout: time.Minute, + } +} + +// Codex/OpenAI 上游改走 HTTP/2 后,池化连接被代理/NAT 静默掐断会成为“死连接”: +// 两端都以为连接存活,请求落上去会挂到 TCP 重传超时(分钟级)才失败。Go 的 +// http2.Transport 默认 ReadIdleTimeout=0(不发健康 PING),无法检测这种死连接。 +// 必须显式启用主动 PING 探测,让死连接被提前剔除,而不是只靠 ResponseHeaderTimeout +// 事后兜底。 +func TestEnableOpenAIHTTP2KeepAlive_EnablesPingHealthCheck(t *testing.T) { + tr := &http.Transport{} + + h2, err := enableOpenAIHTTP2KeepAlive(tr) + require.NoError(t, err) + require.NotNil(t, h2, "必须返回已配置的 *http2.Transport") + + require.Positive(t, h2.ReadIdleTimeout, "必须启用空闲 PING 探测以剔除死连接") + require.Equal(t, openAIHTTP2ReadIdleTimeout, h2.ReadIdleTimeout) + require.Equal(t, openAIHTTP2PingTimeout, h2.PingTimeout, "PING 无响应必须有超时判定") + require.NotNil(t, tr.TLSNextProto["h2"], "http2 必须已挂到底层 http.Transport 上") +} + +// openai_h2 模式构建的 Transport 必须带上 H2 PING 健康探测,从源头剔除死连接。 +func TestBuildUpstreamTransport_OpenAIH2_EnablesPingHealthCheck(t *testing.T) { + tr, err := buildUpstreamTransport(http2KeepAliveTestPoolSettings(), nil, upstreamProtocolModeOpenAIH2) + require.NoError(t, err) + require.True(t, tr.ForceAttemptHTTP2, "openai_h2 必须启用 HTTP/2") + require.NotNil(t, tr.TLSNextProto["h2"], "openai_h2 必须显式配置 http2 以启用 ReadIdleTimeout") +} + +// 非 H2 模式(default/h1)不应因本次改动被误配置:default 走 Go 自动 H2(惰性配置, +// 构建时 TLSNextProto 仍为空),h1 模式显式禁用 H2。避免波及 Claude/Gemini 热路径。 +func TestBuildUpstreamTransport_NonOpenAIH2_NotEagerlyConfigured(t *testing.T) { + tr, err := buildUpstreamTransport(http2KeepAliveTestPoolSettings(), nil, upstreamProtocolModeDefault) + require.NoError(t, err) + require.Nil(t, tr.TLSNextProto["h2"], "default 模式不应在构建期主动配置 http2 keepalive") +} + +// 死连接在经 HTTP 代理(CONNECT 隧道)时最高发,这是带 proxy 账号的真实生产路径: +// 显式 http2 配置须与 Transport.Proxy 同时正确生效,不能相互干扰。 +func TestBuildUpstreamTransport_OpenAIH2_WithHTTPProxy_EnablesKeepAlive(t *testing.T) { + proxyURL, err := url.Parse("http://127.0.0.1:8080") + require.NoError(t, err) + + tr, err := buildUpstreamTransport(http2KeepAliveTestPoolSettings(), proxyURL, upstreamProtocolModeOpenAIH2) + require.NoError(t, err) + require.True(t, tr.ForceAttemptHTTP2) + require.NotNil(t, tr.TLSNextProto["h2"], "经代理的 openai_h2 也必须启用 http2 keepalive") + require.NotNil(t, tr.Proxy, "HTTP 代理仍须通过 Transport.Proxy 生效") +} diff --git a/backend/internal/repository/http_upstream_test.go b/backend/internal/repository/http_upstream_test.go index f331e00589..cb8fe70d5b 100644 --- a/backend/internal/repository/http_upstream_test.go +++ b/backend/internal/repository/http_upstream_test.go @@ -15,6 +15,151 @@ import ( "github.com/stretchr/testify/suite" ) +func TestHTTPUpstreamDoAppliesGrokCLIIdentityBeforeOAuthRoundTrip(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "") + + for _, endpoint := range []string{"responses", "chat/completions"} { + t.Run(endpoint, func(t *testing.T) { + upstream := NewHTTPUpstream(nil) + svc, ok := upstream.(*httpUpstreamService) + require.True(t, ok) + + const accountID int64 = 4084 + isolation := svc.getIsolationMode() + profile := service.HTTPUpstreamProfileDefault + proxyKey := directProxyKey + protocolMode := svc.resolveProtocolMode(profile, proxyKey, nil) + settings := svc.resolvePoolSettings(isolation, 1) + settings = svc.applyProfilePoolSettings(settings, profile) + cacheKey := buildCacheKey(isolation, proxyKey, accountID, protocolMode) + + var capturedHeaders http.Header + svc.clients[cacheKey] = &upstreamClientEntry{ + client: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + capturedHeaders = req.Header.Clone() + statusCode := http.StatusOK + if req.Header.Get("X-XAI-Token-Auth") != "xai-grok-cli" { + statusCode = http.StatusForbidden + } + return &http.Response{ + StatusCode: statusCode, + Header: make(http.Header), + Body: http.NoBody, + Request: req, + }, nil + })}, + proxyKey: proxyKey, + poolKey: buildPoolKey(settings, protocolMode), + protocolMode: protocolMode, + } + + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/"+endpoint, nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + resp, err := svc.Do(req, "", accountID, 1) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + + require.Equal(t, "0.2.93", capturedHeaders.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-cli", capturedHeaders.Get("X-XAI-Token-Auth")) + require.Equal(t, "xai-grok-workspace/0.2.93", capturedHeaders.Get("User-Agent")) + }) + } +} + +func TestApplyGrokCLIProxyHeaders(t *testing.T) { + t.Run("uses pinned stable version for the CLI proxy", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "") + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.93", req.Header.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-cli", req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "xai-grok-workspace/0.2.93", req.Header.Get("User-Agent")) + }) + + t.Run("accepts a valid operator override", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "0.2.95-alpha.1") + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/chat/completions", nil) + require.NoError(t, err) + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.95-alpha.1", req.Header.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-workspace/0.2.95-alpha.1", req.Header.Get("User-Agent")) + }) + + t.Run("rejects an unsafe override", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "0.2.95\r\nX-Injected: true") + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.93", req.Header.Get("x-grok-client-version")) + require.Empty(t, req.Header.Get("X-Injected")) + }) + + t.Run("rejects an override below the supported minimum", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "0.2.92") + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.93", req.Header.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-workspace/0.2.93", req.Header.Get("User-Agent")) + }) + + t.Run("rejects a prerelease override at the minimum version", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "0.2.93-beta.1") + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.93", req.Header.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-workspace/0.2.93", req.Header.Get("User-Agent")) + }) + + for _, version := range []string{ + "0.2.093", + "0.2.94-alpha..1", + "0.3", + "1", + "0.2.95+build.1", + } { + t.Run("rejects invalid semver "+version, func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", version) + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + applyGrokCLIProxyHeaders(req) + + require.Equal(t, "0.2.93", req.Header.Get("x-grok-client-version")) + require.Equal(t, "xai-grok-workspace/0.2.93", req.Header.Get("User-Agent")) + }) + } + + t.Run("leaves direct xAI API requests unchanged", func(t *testing.T) { + t.Setenv("XAI_GROK_CLI_VERSION", "0.2.95") + req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + applyGrokCLIProxyHeaders(req) + + require.Empty(t, req.Header.Get("x-grok-client-version")) + require.Empty(t, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "sub2api-grok/1.0", req.Header.Get("User-Agent")) + }) +} + // HTTPUpstreamSuite HTTP 上游服务测试套件 // 使用 testify/suite 组织测试,支持 SetupTest 初始化 type HTTPUpstreamSuite struct { diff --git a/backend/internal/repository/migrations_runner.go b/backend/internal/repository/migrations_runner.go index 7c045fea74..a071967f65 100644 --- a/backend/internal/repository/migrations_runner.go +++ b/backend/internal/repository/migrations_runner.go @@ -55,6 +55,8 @@ const paymentOrdersOutTradeNoUniqueMigration = "120_enforce_payment_orders_out_t const paymentOrdersOutTradeNoUniqueIndex = "paymentorder_out_trade_no_unique" const schedulerOutboxPendingDedupKeyMigration = "153_scheduler_outbox_pending_dedup_key_index_notx.sql" const schedulerOutboxPendingDedupKeyIndex = "idx_scheduler_outbox_pending_dedup_key" +const latestAPIKeyIPIndexMigration = "174_add_usage_logs_api_key_latest_ip_index_notx.sql" +const latestAPIKeyIPIndex = "idx_usage_logs_api_key_latest_ip" type migrationChecksumCompatibilityRule struct { fileChecksum string @@ -264,6 +266,8 @@ func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name stri return preparePaymentOrdersOutTradeNoUniqueMigration(ctx, db) case schedulerOutboxPendingDedupKeyMigration: return dropInvalidIndexIfPresent(ctx, db, schedulerOutboxPendingDedupKeyIndex) + case latestAPIKeyIPIndexMigration: + return dropInvalidIndexIfPresent(ctx, db, latestAPIKeyIPIndex) default: return nil } diff --git a/backend/internal/repository/migrations_runner_notx_test.go b/backend/internal/repository/migrations_runner_notx_test.go index c9f6a2cdf1..6bb7914b95 100644 --- a/backend/internal/repository/migrations_runner_notx_test.go +++ b/backend/internal/repository/migrations_runner_notx_test.go @@ -116,6 +116,45 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_t_b ON t(b); require.NoError(t, mock.ExpectationsWereMet()) } +func TestApplyMigrationsFS_NonTransactionalMigration_LatestAPIKeyIPIndexDropsInvalidIndexBeforeRetry(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + prepareMigrationsBootstrapExpectations(mock) + mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1"). + WithArgs(latestAPIKeyIPIndexMigration). + WillReturnError(sql.ErrNoRows) + mock.ExpectQuery("SELECT EXISTS \\("). + WithArgs(latestAPIKeyIPIndex). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + mock.ExpectExec("DROP INDEX CONCURRENTLY IF EXISTS idx_usage_logs_api_key_latest_ip"). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip"). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec("INSERT INTO schema_migrations \\(filename, checksum\\) VALUES \\(\\$1, \\$2\\)"). + WithArgs(latestAPIKeyIPIndexMigration, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("SELECT pg_advisory_unlock\\(\\$1\\)"). + WithArgs(migrationsAdvisoryLockID). + WillReturnResult(sqlmock.NewResult(0, 1)) + + fsys := fstest.MapFS{ + latestAPIKeyIPIndexMigration: &fstest.MapFile{ + Data: []byte(` +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip + ON usage_logs (api_key_id, created_at DESC, id DESC) + INCLUDE (ip_address) + WHERE ip_address IS NOT NULL AND ip_address <> ''; +`), + }, + } + + err = applyMigrationsFS(context.Background(), db, fsys) + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestApplyMigrationsFS_PaymentOrdersOutTradeNoUniqueMigration_FailsFastOnDuplicatePrecheck(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) diff --git a/backend/internal/repository/openai_long_context_billing_migration_integration_test.go b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go new file mode 100644 index 0000000000..5f50ed0729 --- /dev/null +++ b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go @@ -0,0 +1,159 @@ +//go:build integration + +package repository + +import ( + "context" + "testing" + + dbmigrations "github.com/Wei-Shaw/sub2api/migrations" + "github.com/stretchr/testify/require" +) + +func TestMigration175EnforcesOpenAILongContextBillingWriteInvariant(t *testing.T) { + tx := testTx(t) + ctx := context.Background() + migrationSQL, err := dbmigrations.FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts; +DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts; +`) + require.NoError(t, err) + + var ordinaryID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-ordinary', 'openai', 'oauth', '{}'::jsonb) +RETURNING id +`).Scan(&ordinaryID)) + + var parentID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-parent', 'openai', 'oauth', '{"openai_long_context_billing_enabled":false}'::jsonb) +RETURNING id +`).Scan(&parentID)) + + var shadowID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra, parent_account_id, quota_dimension) +VALUES ('migration-175-shadow', 'openai', 'oauth', '{}'::jsonb, $1, 'spark') +RETURNING id +`, parentID).Scan(&shadowID)) + + var malformedLegacyID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-malformed-legacy', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb) +RETURNING id +`).Scan(&malformedLegacyID)) + + _, err = tx.ExecContext(ctx, string(migrationSQL)) + require.NoError(t, err) + _, err = tx.ExecContext(ctx, string(migrationSQL)) + require.NoError(t, err) + + var ordinaryEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, ordinaryID).Scan(&ordinaryEnabled)) + require.False(t, ordinaryEnabled) + + var shadowEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.False(t, shadowEnabled) + + var initialShadowOutboxEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&initialShadowOutboxEvents)) + require.Equal(t, 1, initialShadowOutboxEvents) + + var malformedLegacyEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, malformedLegacyID).Scan(&malformedLegacyEnabled)) + require.False(t, malformedLegacyEnabled) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = extra || '{"migration_175_unrelated_update":true}'::jsonb +WHERE id = $1 +`, malformedLegacyID) + require.NoError(t, err) + + _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = '{"legacy_writer_replaced_extra":true}'::jsonb +WHERE id = $1 +`, parentID) + require.NoError(t, err) + var parentEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, parentID).Scan(&parentEnabled)) + require.False(t, parentEnabled) + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.False(t, shadowEnabled) + var preservedOptOutEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&preservedOptOutEvents)) + require.Zero(t, preservedOptOutEvents) + + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-rolling-writer', 'openai', 'oauth', '{}'::jsonb) +RETURNING (extra->>'openai_long_context_billing_enabled')::boolean +`).Scan(&ordinaryEnabled)) + require.False(t, ordinaryEnabled) + + _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = jsonb_set(extra, '{openai_long_context_billing_enabled}', 'true'::jsonb, true) +WHERE id = $1 +`, parentID) + require.NoError(t, err) + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.True(t, shadowEnabled) + + var shadowOutboxEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&shadowOutboxEvents)) + require.Equal(t, 1, shadowOutboxEvents) + + _, err = tx.ExecContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-malformed', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb) +`) + require.ErrorContains(t, err, "openai_long_context_billing_enabled must be a boolean") +} diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go index 2129a451c4..900abcf212 100644 --- a/backend/internal/repository/ops_repo.go +++ b/backend/internal/repository/ops_repo.go @@ -718,6 +718,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser stmt, err := tx.PrepareContext(ctx, pq.CopyIn( "ops_system_logs", "created_at", + "host", "level", "component", "message", @@ -760,6 +761,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser if _, err := stmt.ExecContext( ctx, createdAt.UTC(), + opsNullString(input.Host), level, component, message, @@ -827,6 +829,7 @@ func (r *opsRepository) ListSystemLogs(ctx context.Context, filter *service.OpsS SELECT l.id, l.created_at, + COALESCE(l.host, ''), l.level, COALESCE(l.component, ''), COALESCE(l.message, ''), @@ -859,6 +862,7 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) if err := rows.Scan( &item.ID, &item.CreatedAt, + &item.Host, &item.Level, &item.Component, &item.Message, @@ -1130,6 +1134,11 @@ func buildOpsSystemLogsWhere(filter *service.OpsSystemLogFilter) (string, []any, hasConstraint = true } if filter != nil { + if v := strings.TrimSpace(filter.Host); v != "" { + args = append(args, v) + clauses = append(clauses, "l.host = $"+itoa(len(args))) + hasConstraint = true + } if v := strings.ToLower(strings.TrimSpace(filter.Level)); v != "" { args = append(args, v) clauses = append(clauses, "LOWER(COALESCE(l.level,'')) = $"+itoa(len(args))) @@ -1194,6 +1203,7 @@ func buildOpsSystemLogsCleanupWhere(filter *service.OpsSystemLogCleanupFilter) ( listFilter := &service.OpsSystemLogFilter{ StartTime: filter.StartTime, EndTime: filter.EndTime, + Host: filter.Host, Level: filter.Level, Component: filter.Component, RequestID: filter.RequestID, diff --git a/backend/internal/repository/ops_repo_system_logs_test.go b/backend/internal/repository/ops_repo_system_logs_test.go index 98199f4828..48be3e7256 100644 --- a/backend/internal/repository/ops_repo_system_logs_test.go +++ b/backend/internal/repository/ops_repo_system_logs_test.go @@ -18,6 +18,7 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) { filter := &service.OpsSystemLogFilter{ StartTime: &start, EndTime: &end, + Host: "api-node-1", Level: "warn", Component: "http.access", RequestID: "req-1", @@ -37,8 +38,11 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) { if where == "" { t.Fatalf("where should not be empty") } - if len(args) != 12 { - t.Fatalf("args len = %d, want 12", len(args)) + if len(args) != 13 { + t.Fatalf("args len = %d, want 13", len(args)) + } + if !contains(where, "l.host = $") { + t.Fatalf("where should include host condition: %s", where) } if !contains(where, "COALESCE(l.client_request_id,'') = $") { t.Fatalf("where should include client_request_id condition: %s", where) @@ -68,6 +72,7 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing. userID := int64(9) apiKeyID := int64(10) filter := &service.OpsSystemLogCleanupFilter{ + Host: "api-node-2", ClientRequestID: "creq-9", UserID: &userID, APIKeyID: &apiKeyID, @@ -77,8 +82,11 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing. if !hasConstraint { t.Fatalf("expected hasConstraint=true") } - if len(args) != 3 { - t.Fatalf("args len = %d, want 3", len(args)) + if len(args) != 4 { + t.Fatalf("args len = %d, want 4", len(args)) + } + if !contains(where, "l.host = $") { + t.Fatalf("where should include host condition: %s", where) } if !contains(where, "COALESCE(l.client_request_id,'') = $") { t.Fatalf("where should include client_request_id condition: %s", where) diff --git a/backend/internal/repository/proxy_expiry_integration_test.go b/backend/internal/repository/proxy_expiry_integration_test.go index d0cdc913b6..88f4f1df93 100644 --- a/backend/internal/repository/proxy_expiry_integration_test.go +++ b/backend/internal/repository/proxy_expiry_integration_test.go @@ -4,6 +4,7 @@ package repository import ( "context" + "encoding/json" "testing" "time" @@ -70,6 +71,41 @@ func (s *ProxyExpirySuite) TestSweep_DirectMode() { s.Require().Equal(pid, *origin) } +func (s *ProxyExpirySuite) TestSweep_EnqueuesChangedAccountIDsWithoutFullRebuild() { + past := time.Now().Add(-time.Hour) + firstProxyID := s.mkProxy("p-bulk-first", service.FallbackModeDirect, &past, nil) + secondProxyID := s.mkProxy("p-bulk-second", service.FallbackModeDirect, &past, nil) + firstAccountID := s.mkAccountWithProxy(firstProxyID) + secondAccountID := s.mkAccountWithProxy(secondProxyID) + + changed, err := s.repo.SweepExpiredProxies(s.ctx, time.Now()) + s.Require().NoError(err) + s.Require().EqualValues(2, changed) + + var payloadRaw []byte + err = scanSingleRow(s.ctx, s.tx, ` + SELECT payload + FROM scheduler_outbox + WHERE event_type=$1 + ORDER BY id DESC + LIMIT 1`, []any{service.SchedulerOutboxEventAccountBulkChanged}, &payloadRaw) + s.Require().NoError(err) + + var payload struct { + AccountIDs []int64 `json:"account_ids"` + } + s.Require().NoError(json.Unmarshal(payloadRaw, &payload)) + s.Require().Equal([]int64{firstAccountID, secondAccountID}, payload.AccountIDs) + + var fullRebuildCount int + err = scanSingleRow(s.ctx, s.tx, ` + SELECT COUNT(*) + FROM scheduler_outbox + WHERE event_type=$1`, []any{service.SchedulerOutboxEventFullRebuild}, &fullRebuildCount) + s.Require().NoError(err) + s.Require().Zero(fullRebuildCount) +} + func (s *ProxyExpirySuite) TestSweep_ProxyMode_Healthy() { future := time.Now().Add(24 * time.Hour) past := time.Now().Add(-time.Hour) diff --git a/backend/internal/repository/proxy_expiry_test.go b/backend/internal/repository/proxy_expiry_test.go new file mode 100644 index 0000000000..7313241085 --- /dev/null +++ b/backend/internal/repository/proxy_expiry_test.go @@ -0,0 +1,27 @@ +package repository + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestSortedUniqueAccountIDs(t *testing.T) { + tests := []struct { + name string + input []int64 + want []int64 + }{ + {name: "unsorted duplicates", input: []int64{12, 3, 12, 8, 3}, want: []int64{3, 8, 12}}, + {name: "already sorted", input: []int64{3, 8, 12}, want: []int64{3, 8, 12}}, + {name: "single", input: []int64{3}, want: []int64{3}}, + {name: "empty", input: []int64{}, want: []int64{}}, + {name: "nil", input: nil, want: nil}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, sortedUniqueAccountIDs(tt.input)) + }) + } +} diff --git a/backend/internal/repository/proxy_repo.go b/backend/internal/repository/proxy_repo.go index b34c0cb559..fcb2e53b87 100644 --- a/backend/internal/repository/proxy_repo.go +++ b/backend/internal/repository/proxy_repo.go @@ -489,7 +489,7 @@ func (r *proxyRepository) ListAllForFallback(ctx context.Context) ([]service.Pro // SweepExpiredProxies 扫描到期 active 代理,标记 expired 并按 fallback 策略改写绑定账号的 proxy_id, // 最终触发 scheduler outbox 使 Redis 快照缓存失效。返回受影响的账号行数。 // 原子性边界:每个过期代理的「标记 expired + 改投账号」在各自子事务内原子执行(见 sweepOneExpiredProxy); -// 全部代理处理完后若有账号被改投,再统一 enqueue 一次 full_rebuild 事件——该 enqueue 在子事务之外 +// 全部代理处理完后若有账号被改投,再统一 enqueue 一次 account_bulk_changed 事件——该 enqueue 在子事务之外 // (走 r.sql、失败仅记日志、由调度器周期性 full rebuild 兜底),故「改投 → 失效」整体并非原子。 func (r *proxyRepository) SweepExpiredProxies(ctx context.Context, now time.Time) (int64, error) { // 快照读(事务前):允许脏读不影响正确性,事务内已加锁写。 @@ -503,7 +503,7 @@ func (r *proxyRepository) SweepExpiredProxies(ctx context.Context, now time.Time } var totalChanged int64 - accountsTouched := false + allChangedAccountIDs := make([]int64, 0) for _, p := range all { if p.Status != service.StatusActive || !p.IsExpired(now) { @@ -516,79 +516,116 @@ func (r *proxyRepository) SweepExpiredProxies(ctx context.Context, now time.Time logger.LegacyPrintf("repository.proxy", "[ProxyExpiry] proxy %d expired but fallback chain unresolved (cycle/all-expired); accounts kept", p.ID) } - changed, sweepErr := r.sweepOneExpiredProxy(ctx, p.ID, target, change) + changedAccountIDs, sweepErr := r.sweepOneExpiredProxy(ctx, p.ID, target, change) if sweepErr != nil { return totalChanged, sweepErr } - if changed > 0 { - totalChanged += changed - accountsTouched = true - } + totalChanged += int64(len(changedAccountIDs)) + allChangedAccountIDs = append(allChangedAccountIDs, changedAccountIDs...) } - if accountsTouched { - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventFullRebuild, nil, nil, nil); err != nil { - logger.LegacyPrintf("repository.proxy", "[SchedulerOutbox] enqueue proxy expiry rebuild failed: err=%v", err) + changedAccountIDs := sortedUniqueAccountIDs(allChangedAccountIDs) + if len(changedAccountIDs) > 0 { + // 各代理的改投事务已经提交;这里仅汇总真实被 UPDATE 命中的账号, + // 避免代理到期时用全量重建刷新所有调度分桶。 + payload := map[string]any{"account_ids": changedAccountIDs} + if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil { + logger.LegacyPrintf("repository.proxy", "[SchedulerOutbox] enqueue proxy expiry account changes failed: err=%v", err) } } return totalChanged, nil } +func sortedUniqueAccountIDs(accountIDs []int64) []int64 { + if len(accountIDs) < 2 { + return accountIDs + } + sort.Slice(accountIDs, func(i, j int) bool { return accountIDs[i] < accountIDs[j] }) + write := 1 + for _, accountID := range accountIDs[1:] { + if accountID == accountIDs[write-1] { + continue + } + accountIDs[write] = accountID + write++ + } + return accountIDs[:write] +} + // sweepOneExpiredProxy 在单事务内原子执行:标记代理 expired + 改投绑定账号。 // 若 r.client 已绑定事务(测试注入场景),直接在 r.sql 上执行,由外层事务保证原子性。 -func (r *proxyRepository) sweepOneExpiredProxy(ctx context.Context, proxyID int64, target *int64, change bool) (int64, error) { +func (r *proxyRepository) sweepOneExpiredProxy(ctx context.Context, proxyID int64, target *int64, change bool) ([]int64, error) { // 尝试开启子事务;若 r.client 已是事务 client,则返回 ErrTxStarted,退回使用 r.sql。 tx, txErr := r.client.Tx(ctx) if txErr != nil { if txErr != dbent.ErrTxStarted { - return 0, txErr + return nil, txErr } // 已在外层事务中(集成测试场景),直接用 r.sql 执行 return r.sweepOneExpiredProxyOnExec(ctx, r.sql, proxyID, target, change) } // 使用新事务执行 - var n int64 + var accountIDs []int64 var err error - n, err = r.sweepOneExpiredProxyOnExec(ctx, tx, proxyID, target, change) + accountIDs, err = r.sweepOneExpiredProxyOnExec(ctx, tx, proxyID, target, change) if err != nil { _ = tx.Rollback() - return 0, err + return nil, err } if commitErr := tx.Commit(); commitErr != nil { - return 0, commitErr + return nil, commitErr } - return n, nil + return accountIDs, nil } // sweepOneExpiredProxyOnExec 在给定的 sqlExecutor 上执行:标记 expired + 改投账号。 -func (r *proxyRepository) sweepOneExpiredProxyOnExec(ctx context.Context, exec sqlExecutor, proxyID int64, target *int64, change bool) (int64, error) { +func (r *proxyRepository) sweepOneExpiredProxyOnExec(ctx context.Context, exec sqlExecutor, proxyID int64, target *int64, change bool) ([]int64, error) { if _, err := exec.ExecContext(ctx, `UPDATE proxies SET status=$1, updated_at=NOW() WHERE id=$2 AND deleted_at IS NULL`, service.StatusExpired, proxyID); err != nil { - return 0, err + return nil, err } if !change { - return 0, nil + return nil, nil } var ( - res sql.Result - err error + rows *sql.Rows + err error ) if target == nil { - res, err = exec.ExecContext(ctx, ` + rows, err = exec.QueryContext(ctx, ` UPDATE accounts SET proxy_id=NULL, proxy_fallback_origin_id=$1, updated_at=NOW() - WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL`, proxyID) + WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL + RETURNING id`, proxyID) } else { - res, err = exec.ExecContext(ctx, ` + rows, err = exec.QueryContext(ctx, ` UPDATE accounts SET proxy_id=$2, proxy_fallback_origin_id=$1, updated_at=NOW() - WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL`, proxyID, *target) + WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL + RETURNING id`, proxyID, *target) } if err != nil { - return 0, err + return nil, err } - n, _ := res.RowsAffected() - return n, nil + + // 必须在提交子事务前读完并关闭 RETURNING 结果集,否则连接仍可能处于 busy 状态。 + accountIDs := make([]int64, 0) + for rows.Next() { + var accountID int64 + if err := rows.Scan(&accountID); err != nil { + _ = rows.Close() + return nil, err + } + accountIDs = append(accountIDs, accountID) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + return accountIDs, nil } // CountExpired 返回已过期(status=expired)的代理数量。 diff --git a/backend/internal/repository/redis.go b/backend/internal/repository/redis.go index 2b4ee4e636..0ead4644c1 100644 --- a/backend/internal/repository/redis.go +++ b/backend/internal/repository/redis.go @@ -21,7 +21,11 @@ import ( // 2. MinIdleConns: 保持最小空闲连接,减少冷启动延迟(默认 10) // 3. DialTimeout/ReadTimeout/WriteTimeout: 精确控制各阶段超时 func InitRedis(cfg *config.Config) *redis.Client { - return redis.NewClient(buildRedisOptions(cfg)) + client := redis.NewClient(buildRedisOptions(cfg)) + if cfg.Server.EnableServerTiming { + client.AddHook(serverTimingRedisHook{}) + } + return client } // buildRedisOptions 构建 Redis 连接选项 diff --git a/backend/internal/repository/req_client_pool.go b/backend/internal/repository/req_client_pool.go index 32501f7b19..95ab27ce32 100644 --- a/backend/internal/repository/req_client_pool.go +++ b/backend/internal/repository/req_client_pool.go @@ -2,11 +2,13 @@ package repository import ( "fmt" + "net/http" "strings" "sync" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/imroc/req/v3" ) @@ -57,6 +59,7 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) { if trimmed != "" { client.SetProxyURL(trimmed) } + client = instrumentReqClient(client) actual, _ := sharedReqClients.LoadOrStore(key, client) if c, ok := actual.(*req.Client); ok { @@ -65,6 +68,17 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) { return client, nil } +func instrumentReqClient(client *req.Client) *req.Client { + if client == nil { + return nil + } + client.GetTransport().WrapRoundTripFunc(func(rt http.RoundTripper) req.HttpRoundTripFunc { + timed := servertiming.WrapRoundTripper(rt) + return timed.RoundTrip + }) + return client +} + func buildReqClientKey(opts reqClientOptions) string { return fmt.Sprintf("%s|%s|%t|%t", strings.TrimSpace(opts.ProxyURL), diff --git a/backend/internal/repository/req_client_pool_test.go b/backend/internal/repository/req_client_pool_test.go index 9067d0129f..3a27841c5a 100644 --- a/backend/internal/repository/req_client_pool_test.go +++ b/backend/internal/repository/req_client_pool_test.go @@ -1,12 +1,17 @@ package repository import ( + "context" + "net/http" + "net/http/httptest" "reflect" + "strings" "sync" "testing" "time" "unsafe" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/imroc/req/v3" "github.com/stretchr/testify/require" ) @@ -118,3 +123,20 @@ func TestCreateGeminiReqClient_ForceHTTP2Disabled(t *testing.T) { require.NoError(t, err) require.Equal(t, "", forceHTTPVersion(t, client)) } + +func TestInstrumentReqClientRecordsDependency(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + client := instrumentReqClient(req.C()) + response, err := client.R().SetContext(ctx).Get(server.URL) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, response.StatusCode) + + header := collector.HeaderValue(time.Now(), "bypass") + require.True(t, strings.Contains(header, "dep_http;dur="), header) +} diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index c8e1fe14e0..c68bcf96e6 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "log/slog" "strconv" "time" @@ -163,14 +164,15 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul versionStr := strconv.FormatInt(version, 10) snapshotKey := schedulerSnapshotKey(bucket, versionStr) - if err := c.writeAccounts(ctx, accounts); err != nil { + cacheableAccounts, err := c.writeAccounts(ctx, accounts) + if err != nil { return err } - if len(accounts) > 0 { + if len(cacheableAccounts) > 0 { // 使用序号作为 score,保持数据库返回的排序语义。 - members := make([]redis.Z, 0, len(accounts)) - for idx, account := range accounts { + members := make([]redis.Z, 0, len(cacheableAccounts)) + for idx, account := range cacheableAccounts { members = append(members, redis.Z{ Score: float64(idx), Member: strconv.FormatInt(account.ID, 10), @@ -224,7 +226,14 @@ func (c *schedulerCache) SetAccount(ctx context.Context, account *service.Accoun if account == nil || account.ID <= 0 { return nil } - return c.writeAccounts(ctx, []service.Account{*account}) + cacheableAccounts, err := c.writeAccounts(ctx, []service.Account{*account}) + if err != nil { + return err + } + if len(cacheableAccounts) == 0 { + return c.DeleteAccount(ctx, account.ID) + } + return nil } func (c *schedulerCache) DeleteAccount(ctx context.Context, accountID int64) error { @@ -262,13 +271,14 @@ func (c *schedulerCache) UpdateLastUsed(ctx context.Context, updates map[int64]t return err } account.LastUsedAt = ptrTime(updates[ids[i]]) - updated, err := json.Marshal(account) + updated, metaPayload, err := marshalSchedulerCacheAccount(*account) if err != nil { - return err - } - metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(*account)) - if err != nil { - return err + slog.Warn("scheduler cache removes account with unencodable payload", + "account_id", ids[i], + "error", err, + ) + pipe.Del(ctx, keys[i], schedulerAccountMetaKey(strconv.FormatInt(ids[i], 10))) + continue } pipe.Set(ctx, keys[i], updated, 0) pipe.Set(ctx, schedulerAccountMetaKey(strconv.FormatInt(ids[i], 10)), metaPayload, 0) @@ -359,12 +369,13 @@ func decodeCachedAccount(val any) (*service.Account, error) { return &account, nil } -func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.Account) error { +func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.Account) ([]service.Account, error) { if len(accounts) == 0 { - return nil + return nil, nil } pipe := c.rdb.Pipeline() + cacheableAccounts := make([]service.Account, 0, len(accounts)) pending := 0 flush := func() error { if pending == 0 { @@ -379,27 +390,43 @@ func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.A } for _, account := range accounts { - fullPayload, err := json.Marshal(account) + fullPayload, metaPayload, err := marshalSchedulerCacheAccount(account) if err != nil { - return err - } - metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(account)) - if err != nil { - return err + slog.Warn("scheduler cache skips account with unencodable payload", + "account_id", account.ID, + "error", err, + ) + continue } id := strconv.FormatInt(account.ID, 10) pipe.Set(ctx, schedulerAccountKey(id), fullPayload, 0) pipe.Set(ctx, schedulerAccountMetaKey(id), metaPayload, 0) + cacheableAccounts = append(cacheableAccounts, account) pending++ if pending >= c.writeChunkSize { if err := flush(); err != nil { - return err + return nil, err } } } - return flush() + if err := flush(); err != nil { + return nil, err + } + return cacheableAccounts, nil +} + +func marshalSchedulerCacheAccount(account service.Account) ([]byte, []byte, error) { + fullPayload, err := json.Marshal(account) + if err != nil { + return nil, nil, fmt.Errorf("marshal account: %w", err) + } + metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(account)) + if err != nil { + return nil, nil, fmt.Errorf("marshal account metadata: %w", err) + } + return fullPayload, metaPayload, nil } func (c *schedulerCache) mgetChunked(ctx context.Context, keys []string) ([]any, error) { diff --git a/backend/internal/repository/scheduler_cache_unit_test.go b/backend/internal/repository/scheduler_cache_unit_test.go index 19c4cc4f36..ecca7f3892 100644 --- a/backend/internal/repository/scheduler_cache_unit_test.go +++ b/backend/internal/repository/scheduler_cache_unit_test.go @@ -3,12 +3,78 @@ package repository import ( + "context" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" "github.com/stretchr/testify/require" ) +func newSchedulerCacheUnit(t *testing.T) *schedulerCache { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + cache, ok := newSchedulerCacheWithChunkSizes(rdb, defaultSchedulerSnapshotMGetChunkSize, defaultSchedulerSnapshotWriteChunkSize).(*schedulerCache) + require.True(t, ok) + return cache +} + +func TestSchedulerCacheWriteAccountsSkipsUnencodableTimes(t *testing.T) { + ctx := context.Background() + cache := newSchedulerCacheUnit(t) + invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC) + + cacheable, err := cache.writeAccounts(ctx, []service.Account{ + {ID: 111, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey}, + {ID: 112, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, ExpiresAt: &invalidTime}, + }) + require.NoError(t, err) + require.Len(t, cacheable, 1) + require.Equal(t, int64(111), cacheable[0].ID) + + cached, err := cache.GetAccount(ctx, 111) + require.NoError(t, err) + require.NotNil(t, cached) + + invalid, err := cache.GetAccount(ctx, 112) + require.NoError(t, err) + require.Nil(t, invalid) +} + +func TestSchedulerCacheSetAccountClearsUnencodablePayload(t *testing.T) { + ctx := context.Background() + cache := newSchedulerCacheUnit(t) + + account := service.Account{ID: 113, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey} + require.NoError(t, cache.SetAccount(ctx, &account)) + + invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC) + account.ExpiresAt = &invalidTime + require.NoError(t, cache.SetAccount(ctx, &account)) + + cached, err := cache.GetAccount(ctx, account.ID) + require.NoError(t, err) + require.Nil(t, cached) +} + +func TestSchedulerCacheUpdateLastUsedClearsUnencodablePayload(t *testing.T) { + ctx := context.Background() + cache := newSchedulerCacheUnit(t) + account := service.Account{ID: 114, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey} + require.NoError(t, cache.SetAccount(ctx, &account)) + + invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC) + require.NoError(t, cache.UpdateLastUsed(ctx, map[int64]time.Time{account.ID: invalidTime})) + + cached, err := cache.GetAccount(ctx, account.ID) + require.NoError(t, err) + require.Nil(t, cached) +} + func TestBuildSchedulerMetadataAccount_KeepsOpenAIWSFlags(t *testing.T) { account := service.Account{ ID: 42, diff --git a/backend/internal/repository/scheduler_outbox_repo.go b/backend/internal/repository/scheduler_outbox_repo.go index 59772fbb5a..500f6e850d 100644 --- a/backend/internal/repository/scheduler_outbox_repo.go +++ b/backend/internal/repository/scheduler_outbox_repo.go @@ -93,6 +93,24 @@ func (r *schedulerOutboxRepository) ListAfterAndReleaseDedup(ctx context.Context return events, nil } +func (r *schedulerOutboxRepository) FirstCreatedAtAfter(ctx context.Context, afterID int64) (time.Time, bool, error) { + var createdAt time.Time + err := r.db.QueryRowContext(ctx, ` + SELECT created_at + FROM scheduler_outbox + WHERE id > $1 + ORDER BY id ASC + LIMIT 1 + `, afterID).Scan(&createdAt) + if err == sql.ErrNoRows { + return time.Time{}, false, nil + } + if err != nil { + return time.Time{}, false, err + } + return createdAt, true, nil +} + func (r *schedulerOutboxRepository) MaxID(ctx context.Context) (int64, error) { var maxID int64 if err := r.db.QueryRowContext(ctx, "SELECT COALESCE(MAX(id), 0) FROM scheduler_outbox").Scan(&maxID); err != nil { diff --git a/backend/internal/repository/scheduler_outbox_repo_test.go b/backend/internal/repository/scheduler_outbox_repo_test.go index 619339d207..250014f6ee 100644 --- a/backend/internal/repository/scheduler_outbox_repo_test.go +++ b/backend/internal/repository/scheduler_outbox_repo_test.go @@ -4,11 +4,63 @@ import ( "context" "regexp" "testing" + "time" sqlmock "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/require" ) +func TestSchedulerOutboxRepositoryFirstCreatedAtAfter(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + repo := &schedulerOutboxRepository{db: db} + createdAt := time.Now().UTC().Truncate(time.Microsecond) + const expectedSQL = ` + SELECT created_at + FROM scheduler_outbox + WHERE id > $1 + ORDER BY id ASC + LIMIT 1 + ` + mock.ExpectQuery(regexp.QuoteMeta(expectedSQL)). + WithArgs(int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"created_at"}).AddRow(createdAt)) + + got, ok, err := repo.FirstCreatedAtAfter(context.Background(), 42) + + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, createdAt, got) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestSchedulerOutboxRepositoryFirstCreatedAtAfterReturnsNotFound(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + repo := &schedulerOutboxRepository{db: db} + const expectedSQL = ` + SELECT created_at + FROM scheduler_outbox + WHERE id > $1 + ORDER BY id ASC + LIMIT 1 + ` + mock.ExpectQuery(regexp.QuoteMeta(expectedSQL)). + WithArgs(int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"created_at"})) + + got, ok, err := repo.FirstCreatedAtAfter(context.Background(), 42) + + require.NoError(t, err) + require.False(t, ok) + require.True(t, got.IsZero()) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestSchedulerOutboxRepositoryDeleteConsumedUpToUsesBoundedCTE(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) diff --git a/backend/internal/repository/server_timing_redis.go b/backend/internal/repository/server_timing_redis.go new file mode 100644 index 0000000000..dba35450de --- /dev/null +++ b/backend/internal/repository/server_timing_redis.go @@ -0,0 +1,39 @@ +package repository + +import ( + "context" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/redis/go-redis/v9" +) + +type serverTimingRedisHook struct{} + +func (serverTimingRedisHook) DialHook(next redis.DialHook) redis.DialHook { + return next +} + +func (serverTimingRedisHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, cmd redis.Cmder) error { + if !servertiming.Active(ctx) { + return next(ctx, cmd) + } + startedAt := time.Now() + err := next(ctx, cmd) + servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), 1) + return err + } +} + +func (serverTimingRedisHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return func(ctx context.Context, cmds []redis.Cmder) error { + if !servertiming.Active(ctx) { + return next(ctx, cmds) + } + startedAt := time.Now() + err := next(ctx, cmds) + servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), len(cmds)) + return err + } +} diff --git a/backend/internal/repository/server_timing_redis_test.go b/backend/internal/repository/server_timing_redis_test.go new file mode 100644 index 0000000000..d1ae47e3b0 --- /dev/null +++ b/backend/internal/repository/server_timing_redis_test.go @@ -0,0 +1,63 @@ +package repository + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/redis/go-redis/v9" +) + +func TestServerTimingRedisHookRecordsCommands(t *testing.T) { + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + hook := serverTimingRedisHook{} + + process := hook.ProcessHook(func(context.Context, redis.Cmder) error { + time.Sleep(time.Millisecond) + return errors.New("redis failure") + }) + if err := process(ctx, redis.NewStringCmd(ctx, "get", "sensitive-key")); err == nil { + t.Fatal("ProcessHook did not return the underlying error") + } + + pipeline := hook.ProcessPipelineHook(func(context.Context, []redis.Cmder) error { + time.Sleep(time.Millisecond) + return nil + }) + commands := []redis.Cmder{ + redis.NewStringCmd(ctx, "get", "first-secret"), + redis.NewStringCmd(ctx, "get", "second-secret"), + redis.NewStatusCmd(ctx, "set", "third-secret", "value"), + } + if err := pipeline(ctx, commands); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `commands=4`) { + t.Fatalf("header %q does not report one command and a three-command pipeline", header) + } + if strings.Contains(header, "secret") || strings.Contains(header, "get") { + t.Fatalf("Redis command details leaked into header: %q", header) + } +} + +func TestServerTimingRedisHookSkipsInactiveContext(t *testing.T) { + called := false + hook := serverTimingRedisHook{} + process := hook.ProcessHook(func(context.Context, redis.Cmder) error { + called = true + return nil + }) + ctx := context.Background() + if err := process(ctx, redis.NewStringCmd(ctx, "ping")); err != nil { + t.Fatal(err) + } + if !called { + t.Fatal("inactive Redis command did not reach the next hook") + } +} diff --git a/backend/internal/repository/server_timing_sql.go b/backend/internal/repository/server_timing_sql.go new file mode 100644 index 0000000000..062663f08b --- /dev/null +++ b/backend/internal/repository/server_timing_sql.go @@ -0,0 +1,311 @@ +package repository + +import ( + "context" + "database/sql/driver" + "errors" + "io" + "reflect" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" +) + +type serverTimingConnector struct { + base driver.Connector +} + +func newServerTimingConnector(base driver.Connector) driver.Connector { + return &serverTimingConnector{base: base} +} + +func (c *serverTimingConnector) Connect(ctx context.Context) (driver.Conn, error) { + startedAt := time.Now() + conn, err := c.base.Connect(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + if err != nil { + return nil, err + } + return &serverTimingConn{Conn: conn}, nil +} + +func (c *serverTimingConnector) Driver() driver.Driver { + return c.base.Driver() +} + +type serverTimingConn struct { + driver.Conn +} + +func (c *serverTimingConn) Prepare(query string) (driver.Stmt, error) { + stmt, err := c.Conn.Prepare(query) + if err != nil { + return nil, err + } + return &serverTimingStmt{Stmt: stmt}, nil +} + +func (c *serverTimingConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) { + startedAt := time.Now() + var ( + stmt driver.Stmt + err error + ) + if preparer, ok := c.Conn.(driver.ConnPrepareContext); ok { + stmt, err = preparer.PrepareContext(ctx, query) + } else { + stmt, err = c.Conn.Prepare(query) + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil { + return nil, err + } + return &serverTimingStmt{Stmt: stmt}, nil +} + +func (c *serverTimingConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) { + execer, ok := c.Conn.(driver.ExecerContext) + if !ok { + return nil, driver.ErrSkip + } + startedAt := time.Now() + result, err := execer.ExecContext(ctx, query, args) + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + return result, err +} + +func (c *serverTimingConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) { + queryer, ok := c.Conn.(driver.QueryerContext) + if !ok { + return nil, driver.ErrSkip + } + startedAt := time.Now() + rows, err := queryer.QueryContext(ctx, query, args) + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil || rows == nil { + return rows, err + } + return newServerTimingRows(ctx, rows), nil +} + +func (c *serverTimingConn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) { + startedAt := time.Now() + var ( + tx driver.Tx + err error + ) + if beginner, ok := c.Conn.(driver.ConnBeginTx); ok { + tx, err = beginner.BeginTx(ctx, opts) + } else { + if opts.Isolation != driver.IsolationLevel(0) { + return nil, errors.New("driver does not support non-default isolation") + } + if opts.ReadOnly { + return nil, errors.New("driver does not support read-only transactions") + } + // The wrapper exposes ConnBeginTx, so it must retain database/sql's + // legacy fallback for drivers that only implement Conn.Begin. + tx, err = c.Conn.Begin() //nolint:staticcheck // Required driver compatibility fallback. + } + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + if err != nil || tx == nil { + return tx, err + } + return &serverTimingTx{Tx: tx, ctx: ctx}, nil +} + +func (c *serverTimingConn) Ping(ctx context.Context) error { + if pinger, ok := c.Conn.(driver.Pinger); ok { + startedAt := time.Now() + err := pinger.Ping(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err + } + return nil +} + +func (c *serverTimingConn) ResetSession(ctx context.Context) error { + if resetter, ok := c.Conn.(driver.SessionResetter); ok { + startedAt := time.Now() + err := resetter.ResetSession(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err + } + return nil +} + +func (c *serverTimingConn) IsValid() bool { + if validator, ok := c.Conn.(driver.Validator); ok { + return validator.IsValid() + } + return true +} + +func (c *serverTimingConn) CheckNamedValue(value *driver.NamedValue) error { + if checker, ok := c.Conn.(driver.NamedValueChecker); ok { + return checker.CheckNamedValue(value) + } + return driver.ErrSkip +} + +type serverTimingStmt struct { + driver.Stmt +} + +func (s *serverTimingStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) { + startedAt := time.Now() + var ( + result driver.Result + err error + ) + if execer, ok := s.Stmt.(driver.StmtExecContext); ok { + result, err = execer.ExecContext(ctx, args) + } else { + var values []driver.Value + values, err = namedValues(args) + if err == nil { + // The wrapper exposes StmtExecContext and must preserve the fallback + // database/sql would use for a legacy driver statement. + result, err = s.Stmt.Exec(values) //nolint:staticcheck // Required driver compatibility fallback. + } + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + return result, err +} + +func (s *serverTimingStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) { + startedAt := time.Now() + var ( + rows driver.Rows + err error + ) + if queryer, ok := s.Stmt.(driver.StmtQueryContext); ok { + rows, err = queryer.QueryContext(ctx, args) + } else { + var values []driver.Value + values, err = namedValues(args) + if err == nil { + // The wrapper exposes StmtQueryContext and must preserve the fallback + // database/sql would use for a legacy driver statement. + rows, err = s.Stmt.Query(values) //nolint:staticcheck // Required driver compatibility fallback. + } + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil || rows == nil { + return rows, err + } + return newServerTimingRows(ctx, rows), nil +} + +func (s *serverTimingStmt) CheckNamedValue(value *driver.NamedValue) error { + if checker, ok := s.Stmt.(driver.NamedValueChecker); ok { + return checker.CheckNamedValue(value) + } + return driver.ErrSkip +} + +func namedValues(args []driver.NamedValue) ([]driver.Value, error) { + values := make([]driver.Value, len(args)) + for i, arg := range args { + if arg.Name != "" { + return nil, errors.New("named parameters are not supported") + } + values[i] = arg.Value + } + return values, nil +} + +type serverTimingRows struct { + driver.Rows + ctx context.Context +} + +func newServerTimingRows(ctx context.Context, rows driver.Rows) *serverTimingRows { + return &serverTimingRows{Rows: rows, ctx: ctx} +} + +func (r *serverTimingRows) Close() error { + startedAt := time.Now() + err := r.Rows.Close() + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) Next(dest []driver.Value) error { + startedAt := time.Now() + err := r.Rows.Next(dest) + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) HasNextResultSet() bool { + if rows, ok := r.Rows.(driver.RowsNextResultSet); ok { + return rows.HasNextResultSet() + } + return false +} + +func (r *serverTimingRows) NextResultSet() error { + rows, ok := r.Rows.(driver.RowsNextResultSet) + if !ok { + return io.EOF + } + startedAt := time.Now() + err := rows.NextResultSet() + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) ColumnTypeScanType(index int) reflect.Type { + if rows, ok := r.Rows.(driver.RowsColumnTypeScanType); ok { + return rows.ColumnTypeScanType(index) + } + return reflect.TypeOf(new(any)).Elem() +} + +func (r *serverTimingRows) ColumnTypeDatabaseTypeName(index int) string { + if rows, ok := r.Rows.(driver.RowsColumnTypeDatabaseTypeName); ok { + return rows.ColumnTypeDatabaseTypeName(index) + } + return "" +} + +func (r *serverTimingRows) ColumnTypeLength(index int) (int64, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypeLength); ok { + return rows.ColumnTypeLength(index) + } + return 0, false +} + +func (r *serverTimingRows) ColumnTypeNullable(index int) (bool, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypeNullable); ok { + return rows.ColumnTypeNullable(index) + } + return false, false +} + +func (r *serverTimingRows) ColumnTypePrecisionScale(index int) (int64, int64, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypePrecisionScale); ok { + return rows.ColumnTypePrecisionScale(index) + } + return 0, 0, false +} + +type serverTimingTx struct { + driver.Tx + ctx context.Context +} + +func (t *serverTimingTx) Commit() error { + startedAt := time.Now() + err := t.Tx.Commit() + servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (t *serverTimingTx) Rollback() error { + startedAt := time.Now() + err := t.Tx.Rollback() + servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} diff --git a/backend/internal/repository/server_timing_sql_test.go b/backend/internal/repository/server_timing_sql_test.go new file mode 100644 index 0000000000..3a8bbbe03e --- /dev/null +++ b/backend/internal/repository/server_timing_sql_test.go @@ -0,0 +1,258 @@ +package repository + +import ( + "context" + "database/sql/driver" + "io" + "regexp" + "strconv" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" +) + +const fakeDriverDelay = 2 * time.Millisecond + +type timingFakeDriver struct{} + +func (timingFakeDriver) Open(string) (driver.Conn, error) { return newTimingFakeConn(), nil } + +type timingFakeConnector struct { + conn driver.Conn +} + +func (c timingFakeConnector) Connect(context.Context) (driver.Conn, error) { + time.Sleep(fakeDriverDelay) + return c.conn, nil +} + +func (timingFakeConnector) Driver() driver.Driver { return timingFakeDriver{} } + +type timingFakeConn struct{} + +func newTimingFakeConn() *timingFakeConn { return &timingFakeConn{} } + +func (c *timingFakeConn) Prepare(string) (driver.Stmt, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeStmt{}, nil +} + +func (c *timingFakeConn) PrepareContext(context.Context, string) (driver.Stmt, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeStmt{}, nil +} + +func (c *timingFakeConn) Close() error { return nil } + +func (c *timingFakeConn) Begin() (driver.Tx, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeTx{}, nil +} + +func (c *timingFakeConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeTx{}, nil +} + +func (c *timingFakeConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (c *timingFakeConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +func (c *timingFakeConn) Ping(context.Context) error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (c *timingFakeConn) ResetSession(context.Context) error { + time.Sleep(fakeDriverDelay) + return nil +} + +type timingFakeStmt struct{} + +func (s *timingFakeStmt) Close() error { return nil } +func (s *timingFakeStmt) NumInput() int { return -1 } + +func (s *timingFakeStmt) Exec([]driver.Value) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (s *timingFakeStmt) Query([]driver.Value) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +func (s *timingFakeStmt) ExecContext(context.Context, []driver.NamedValue) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (s *timingFakeStmt) QueryContext(context.Context, []driver.NamedValue) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +type timingFakeRows struct { + values [][]driver.Value + index int +} + +func (r *timingFakeRows) Columns() []string { return []string{"value"} } + +func (r *timingFakeRows) Close() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (r *timingFakeRows) Next(dest []driver.Value) error { + time.Sleep(fakeDriverDelay) + if r.index >= len(r.values) { + return io.EOF + } + copy(dest, r.values[r.index]) + r.index++ + return nil +} + +type timingFakeTx struct{} + +func (t *timingFakeTx) Commit() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (t *timingFakeTx) Rollback() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func metricDuration(t *testing.T, header, metric string) float64 { + t.Helper() + re := regexp.MustCompile(`(?:^|, )` + regexp.QuoteMeta(metric) + `;dur=([0-9]+(?:\.[0-9]+)?)`) + match := re.FindStringSubmatch(header) + if len(match) != 2 { + t.Fatalf("metric %q missing from header %q", metric, header) + } + value, err := strconv.ParseFloat(match[1], 64) + if err != nil { + t.Fatalf("parse %s duration: %v", metric, err) + } + return value +} + +func TestServerTimingConnectorRecordsDriverCallsWithoutRowLifetime(t *testing.T) { + startedAt := time.Now() + collector := servertiming.New(startedAt) + ctx := servertiming.WithCollector(context.Background(), collector) + + wrapped := newServerTimingConnector(timingFakeConnector{conn: newTimingFakeConn()}) + rawConn, err := wrapped.Connect(ctx) + if err != nil { + t.Fatal(err) + } + conn, ok := rawConn.(*serverTimingConn) + if !ok { + t.Fatalf("Connect() returned %T, want *serverTimingConn", rawConn) + } + + if _, err := conn.ExecContext(ctx, "sensitive update", nil); err != nil { + t.Fatal(err) + } + rows, err := conn.QueryContext(ctx, "sensitive select", nil) + if err != nil { + t.Fatal(err) + } + values := make([]driver.Value, 1) + if err := rows.Next(values); err != nil { + t.Fatal(err) + } + + // Application work between row reads must remain app time. + time.Sleep(30 * time.Millisecond) + if err := rows.Next(values); err != io.EOF { + t.Fatalf("rows.Next() = %v, want EOF", err) + } + if err := rows.Close(); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `queries=2`) { + t.Fatalf("header %q does not report two SQL operations", header) + } + if strings.Contains(header, "sensitive") { + t.Fatalf("SQL text leaked into header: %q", header) + } + if app, db := metricDuration(t, header, "app"), metricDuration(t, header, "db"); app <= db { + t.Fatalf("row processing gap was counted as DB time: app=%.1fms db=%.1fms header=%q", app, db, header) + } +} + +func TestServerTimingPreparedStatementsAndTransactions(t *testing.T) { + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + conn := &serverTimingConn{Conn: newTimingFakeConn()} + + stmt, err := conn.PrepareContext(ctx, "prepare sensitive statement") + if err != nil { + t.Fatal(err) + } + timedStmt, ok := stmt.(*serverTimingStmt) + if !ok { + t.Fatalf("PrepareContext() returned %T, want *serverTimingStmt", stmt) + } + if _, err := timedStmt.ExecContext(ctx, nil); err != nil { + t.Fatal(err) + } + rows, err := timedStmt.QueryContext(ctx, nil) + if err != nil { + t.Fatal(err) + } + if err := rows.Close(); err != nil { + t.Fatal(err) + } + + tx, err := conn.BeginTx(ctx, driver.TxOptions{}) + if err != nil { + t.Fatal(err) + } + if err := tx.Commit(); err != nil { + t.Fatal(err) + } + if err := conn.Ping(ctx); err != nil { + t.Fatal(err) + } + if err := conn.ResetSession(ctx); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `queries=3`) { + t.Fatalf("header %q does not report prepare, exec, and query operations", header) + } + if metricDuration(t, header, "db") <= 0 { + t.Fatalf("DB duration was not recorded: %q", header) + } +} + +func TestNamedValuesRejectNamedParameters(t *testing.T) { + if _, err := namedValues([]driver.NamedValue{{Name: "secret", Value: 1}}); err == nil { + t.Fatal("namedValues accepted a named parameter") + } + values, err := namedValues([]driver.NamedValue{{Ordinal: 1, Value: "value"}}) + if err != nil { + t.Fatal(err) + } + if len(values) != 1 || values[0] != "value" { + t.Fatalf("namedValues() = %#v", values) + } +} diff --git a/backend/internal/repository/usage_log_repo_insert.go b/backend/internal/repository/usage_log_repo_insert.go index dfd8969512..ec09b308a0 100644 --- a/backend/internal/repository/usage_log_repo_insert.go +++ b/backend/internal/repository/usage_log_repo_insert.go @@ -71,6 +71,7 @@ var usageLogInsertArgTypes = [...]string{ "text", // inbound_endpoint "text", // upstream_endpoint "boolean", // cache_ttl_overridden + "boolean", // long_context_billing_applied "bigint", // channel_id "text", // model_mapping_chain "text", // billing_tier @@ -263,6 +264,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -275,7 +277,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53 + $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54 ) ON CONFLICT (request_id, api_key_id) DO NOTHING RETURNING id, created_at @@ -714,6 +716,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -722,7 +725,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage created_at ) AS (VALUES `) - args := make([]any, 0, len(keys)*53) + args := make([]any, 0, len(keys)*54) argPos := 1 for idx, key := range keys { if idx > 0 { @@ -798,6 +801,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -853,6 +857,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -948,6 +953,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -956,7 +962,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( created_at ) AS (VALUES `) - args := make([]any, 0, len(preparedList)*53) + args := make([]any, 0, len(preparedList)*54) argPos := 1 for idx, prepared := range preparedList { if idx > 0 { @@ -1029,6 +1035,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1084,6 +1091,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1147,6 +1155,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1159,7 +1168,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53 + $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54 ) ON CONFLICT (request_id, api_key_id) DO NOTHING `, prepared.args...) @@ -1264,6 +1273,7 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared { inboundEndpoint, upstreamEndpoint, log.CacheTTLOverridden, + log.LongContextBillingApplied, channelID, modelMappingChain, billingTier, diff --git a/backend/internal/repository/usage_log_repo_query.go b/backend/internal/repository/usage_log_repo_query.go index c178429bab..1fdedd8665 100644 --- a/backend/internal/repository/usage_log_repo_query.go +++ b/backend/internal/repository/usage_log_repo_query.go @@ -19,7 +19,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" ) -const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at" +const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at" func (r *usageLogRepository) GetByID(ctx context.Context, id int64) (log *service.UsageLog, err error) { query := "SELECT " + usageLogSelectColumns + " FROM usage_logs WHERE id = $1" @@ -425,60 +425,61 @@ func (r *usageLogRepository) loadSubscriptions(ctx context.Context, ids []int64) func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, error) { var ( - id int64 - userID int64 - apiKeyID int64 - accountID int64 - requestID sql.NullString - model string - requestedModel sql.NullString - upstreamModel sql.NullString - groupID sql.NullInt64 - subscriptionID sql.NullInt64 - inputTokens int - outputTokens int - cacheCreationTokens int - cacheReadTokens int - cacheCreation5m int - cacheCreation1h int - imageOutputTokens int - imageOutputCost float64 - inputCost float64 - outputCost float64 - cacheCreationCost float64 - cacheReadCost float64 - totalCost float64 - actualCost float64 - rateMultiplier float64 - accountRateMultiplier sql.NullFloat64 - billingType int16 - requestTypeRaw int16 - stream bool - openaiWSMode bool - durationMs sql.NullInt64 - firstTokenMs sql.NullInt64 - userAgent sql.NullString - ipAddress sql.NullString - imageCount int - imageSize sql.NullString - imageInputSize sql.NullString - imageOutputSize sql.NullString - imageSizeSource sql.NullString - imageSizeBreakdown sql.NullString - videoCount int - videoResolution sql.NullString - videoDurationSeconds sql.NullInt64 - serviceTier sql.NullString - reasoningEffort sql.NullString - inboundEndpoint sql.NullString - upstreamEndpoint sql.NullString - cacheTTLOverridden bool - channelID sql.NullInt64 - modelMappingChain sql.NullString - billingTier sql.NullString - billingMode sql.NullString - accountStatsCost sql.NullFloat64 - createdAt time.Time + id int64 + userID int64 + apiKeyID int64 + accountID int64 + requestID sql.NullString + model string + requestedModel sql.NullString + upstreamModel sql.NullString + groupID sql.NullInt64 + subscriptionID sql.NullInt64 + inputTokens int + outputTokens int + cacheCreationTokens int + cacheReadTokens int + cacheCreation5m int + cacheCreation1h int + imageOutputTokens int + imageOutputCost float64 + inputCost float64 + outputCost float64 + cacheCreationCost float64 + cacheReadCost float64 + totalCost float64 + actualCost float64 + rateMultiplier float64 + accountRateMultiplier sql.NullFloat64 + billingType int16 + requestTypeRaw int16 + stream bool + openaiWSMode bool + durationMs sql.NullInt64 + firstTokenMs sql.NullInt64 + userAgent sql.NullString + ipAddress sql.NullString + imageCount int + imageSize sql.NullString + imageInputSize sql.NullString + imageOutputSize sql.NullString + imageSizeSource sql.NullString + imageSizeBreakdown sql.NullString + videoCount int + videoResolution sql.NullString + videoDurationSeconds sql.NullInt64 + serviceTier sql.NullString + reasoningEffort sql.NullString + inboundEndpoint sql.NullString + upstreamEndpoint sql.NullString + cacheTTLOverridden bool + longContextBillingApplied bool + channelID sql.NullInt64 + modelMappingChain sql.NullString + billingTier sql.NullString + billingMode sql.NullString + accountStatsCost sql.NullFloat64 + createdAt time.Time ) if err := scanner.Scan( @@ -530,6 +531,7 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e &inboundEndpoint, &upstreamEndpoint, &cacheTTLOverridden, + &longContextBillingApplied, &channelID, &modelMappingChain, &billingTier, @@ -541,34 +543,35 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e } log := &service.UsageLog{ - ID: id, - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - Model: model, - RequestedModel: coalesceTrimmedString(requestedModel, model), - InputTokens: inputTokens, - OutputTokens: outputTokens, - CacheCreationTokens: cacheCreationTokens, - CacheReadTokens: cacheReadTokens, - CacheCreation5mTokens: cacheCreation5m, - CacheCreation1hTokens: cacheCreation1h, - ImageOutputTokens: imageOutputTokens, - ImageOutputCost: imageOutputCost, - InputCost: inputCost, - OutputCost: outputCost, - CacheCreationCost: cacheCreationCost, - CacheReadCost: cacheReadCost, - TotalCost: totalCost, - ActualCost: actualCost, - RateMultiplier: rateMultiplier, - AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier), - BillingType: int8(billingType), - RequestType: service.RequestTypeFromInt16(requestTypeRaw), - ImageCount: imageCount, - VideoCount: videoCount, - CacheTTLOverridden: cacheTTLOverridden, - CreatedAt: createdAt, + ID: id, + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + Model: model, + RequestedModel: coalesceTrimmedString(requestedModel, model), + InputTokens: inputTokens, + OutputTokens: outputTokens, + CacheCreationTokens: cacheCreationTokens, + CacheReadTokens: cacheReadTokens, + CacheCreation5mTokens: cacheCreation5m, + CacheCreation1hTokens: cacheCreation1h, + ImageOutputTokens: imageOutputTokens, + ImageOutputCost: imageOutputCost, + InputCost: inputCost, + OutputCost: outputCost, + CacheCreationCost: cacheCreationCost, + CacheReadCost: cacheReadCost, + TotalCost: totalCost, + ActualCost: actualCost, + RateMultiplier: rateMultiplier, + AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier), + BillingType: int8(billingType), + RequestType: service.RequestTypeFromInt16(requestTypeRaw), + ImageCount: imageCount, + VideoCount: videoCount, + CacheTTLOverridden: cacheTTLOverridden, + LongContextBillingApplied: longContextBillingApplied, + CreatedAt: createdAt, } // 先回填 legacy 字段,再基于 legacy + request_type 计算最终请求类型,保证历史数据兼容。 log.Stream = stream diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index c32ad2b63f..052c319183 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -88,6 +88,7 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) { sqlmock.AnyArg(), // inbound_endpoint sqlmock.AnyArg(), // upstream_endpoint log.CacheTTLOverridden, + log.LongContextBillingApplied, sqlmock.AnyArg(), // channel_id sqlmock.AnyArg(), // model_mapping_chain sqlmock.AnyArg(), // billing_tier @@ -174,6 +175,7 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) { sqlmock.AnyArg(), sqlmock.AnyArg(), log.CacheTTLOverridden, + log.LongContextBillingApplied, sqlmock.AnyArg(), // channel_id sqlmock.AnyArg(), // model_mapping_chain sqlmock.AnyArg(), // billing_tier @@ -813,6 +815,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, sql.NullString{}, sql.NullString{}, @@ -884,6 +887,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier @@ -939,6 +943,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier @@ -994,6 +999,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index d260afe738..a5e3fde155 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -366,6 +366,7 @@ func TestAPIContracts(t *testing.T) { "video_price_480p": null, "video_price_720p": null, "video_price_1080p": null, + "web_search_price_per_call": null, "allow_image_generation": false, "allow_batch_image_generation": false, "batch_image_discount_multiplier": 0, @@ -593,6 +594,7 @@ func TestAPIContracts(t *testing.T) { "total_cost": 0.5, "actual_cost": 0.5, "rate_multiplier": 1, + "long_context_billing_applied": false, "billing_type": 0, "stream": true, "duration_ms": 100, diff --git a/backend/internal/server/middleware/cors.go b/backend/internal/server/middleware/cors.go index 03d5d025de..0283d53115 100644 --- a/backend/internal/server/middleware/cors.go +++ b/backend/internal/server/middleware/cors.go @@ -52,7 +52,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc { } allowHeaders := []string{ "Content-Type", "Content-Length", "Accept-Encoding", "X-CSRF-Token", "Authorization", - "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key", + "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key", "X-Admin-UI-Request", } // OpenAI Node SDK 会发送 x-stainless-* 请求头,需在 CORS 中显式放行。 openAIProperties := []string{ @@ -83,7 +83,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc { } c.Writer.Header().Set("Access-Control-Allow-Headers", allowHeadersValue) c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH") - c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag") + c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag, Server-Timing") c.Writer.Header().Set("Access-Control-Max-Age", "86400") } // 处理预检请求 diff --git a/backend/internal/server/middleware/cors_test.go b/backend/internal/server/middleware/cors_test.go index 6d0bea3608..6a61f696df 100644 --- a/backend/internal/server/middleware/cors_test.go +++ b/backend/internal/server/middleware/cors_test.go @@ -103,8 +103,10 @@ func TestCORS_AllowedOrigin_HasAllowHeaders(t *testing.T) { // 应设置 Allow-Headers、Allow-Methods 和 Max-Age assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Headers"), "允许的 origin 应收到 Allow-Headers") + assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "X-Admin-UI-Request") assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods"), "允许的 origin 应收到 Allow-Methods") + assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "Server-Timing") assert.Equal(t, "86400", w.Header().Get("Access-Control-Max-Age"), "允许的 origin 应收到 Max-Age=86400") assert.Equal(t, "https://allowed.example.com", w.Header().Get("Access-Control-Allow-Origin"), diff --git a/backend/internal/server/middleware/server_timing.go b/backend/internal/server/middleware/server_timing.go new file mode 100644 index 0000000000..2bb21071e0 --- /dev/null +++ b/backend/internal/server/middleware/server_timing.go @@ -0,0 +1,132 @@ +package middleware + +import ( + "net/http" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/gin-gonic/gin" +) + +const ( + snapshotCacheHeader = "X-Snapshot-Cache" + usageCacheHeader = "X-Usage-Stats-Cache" +) + +type serverTimingResponseWriter struct { + gin.ResponseWriter + context *gin.Context + once sync.Once +} + +func (w *serverTimingResponseWriter) Unwrap() http.ResponseWriter { + return w.ResponseWriter +} + +// ServerTiming collects timing only for requests made by the Admin web UI. +func ServerTiming(enabled bool) gin.HandlerFunc { + if !enabled { + return func(c *gin.Context) { + c.Next() + } + } + return func(c *gin.Context) { + if !isAdminUIRequest(c) || c.Request == nil { + c.Next() + return + } + + collector := servertiming.New(time.Now()) + c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector)) + writer := &serverTimingResponseWriter{ + ResponseWriter: c.Writer, + context: c, + } + c.Writer = writer + c.Next() + writer.finalize() + } +} + +func (w *serverTimingResponseWriter) WriteHeader(statusCode int) { + w.ResponseWriter.WriteHeader(statusCode) +} + +func (w *serverTimingResponseWriter) WriteHeaderNow() { + w.finalize() + w.ResponseWriter.WriteHeaderNow() +} + +func (w *serverTimingResponseWriter) Write(data []byte) (int, error) { + w.finalize() + return w.ResponseWriter.Write(data) +} + +func (w *serverTimingResponseWriter) WriteString(data string) (int, error) { + w.finalize() + return w.ResponseWriter.WriteString(data) +} + +func (w *serverTimingResponseWriter) Flush() { + w.finalize() + w.ResponseWriter.Flush() +} + +func (w *serverTimingResponseWriter) finalize() { + if w == nil { + return + } + w.once.Do(func() { + if value := ServerTimingHeaderValue(w.context); value != "" { + w.ResponseWriter.Header().Set(servertiming.HeaderName, value) + } + }) +} + +// ServerTimingHeaderValue returns a timing value only for an authenticated admin. +func ServerTimingHeaderValue(c *gin.Context) string { + if c == nil || c.Request == nil { + return "" + } + role, ok := GetUserRoleFromContext(c) + if !ok || role != "admin" { + return "" + } + return servertiming.HeaderValue(c.Request.Context(), time.Now(), responseCacheStatus(c.Writer.Header())) +} + +// ServerTimingResponseHeader builds the extra header map required by WebSocket upgrades. +func ServerTimingResponseHeader(c *gin.Context) http.Header { + value := ServerTimingHeaderValue(c) + if value == "" { + return nil + } + return http.Header{servertiming.HeaderName: []string{value}} +} + +func isAdminUIRequest(c *gin.Context) bool { + if c == nil || c.Request == nil || c.Request.URL == nil { + return false + } + if strings.TrimSpace(c.GetHeader(servertiming.AdminUIHeader)) == "1" { + return true + } + path := strings.TrimSpace(c.Request.URL.Path) + return path == "/api/v1/admin" || strings.HasPrefix(path, "/api/v1/admin/") +} + +func responseCacheStatus(header http.Header) string { + for _, name := range []string{snapshotCacheHeader, usageCacheHeader} { + switch strings.ToLower(strings.TrimSpace(header.Get(name))) { + case "hit": + return "hit" + case "miss": + return "miss" + case "bypass": + return "bypass" + } + } + return "bypass" +} diff --git a/backend/internal/server/middleware/server_timing_test.go b/backend/internal/server/middleware/server_timing_test.go new file mode 100644 index 0000000000..c064840ece --- /dev/null +++ b/backend/internal/server/middleware/server_timing_test.go @@ -0,0 +1,188 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/gin-gonic/gin" +) + +func runServerTimingRequest( + t *testing.T, + enabled bool, + path string, + marker string, + role string, + handler gin.HandlerFunc, +) *httptest.ResponseRecorder { + t.Helper() + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(ServerTiming(enabled)) + engine.Any("/*path", func(c *gin.Context) { + if role != "" { + c.Set(string(ContextKeyUserRole), role) + } + handler(c) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, path, nil) + if marker != "" { + request.Header.Set(servertiming.AdminUIHeader, marker) + } + engine.ServeHTTP(recorder, request) + return recorder +} + +func TestServerTimingScopesAndRoleGate(t *testing.T) { + tests := []struct { + name string + enabled bool + path string + marker string + role string + wantHeader bool + }{ + {name: "disabled", enabled: false, path: "/api/v1/admin/users", role: "admin"}, + {name: "admin API path", enabled: true, path: "/api/v1/admin/users", role: "admin", wantHeader: true}, + {name: "shared API marked by admin UI", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "admin", wantHeader: true}, + {name: "non admin role", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "user"}, + {name: "unauthenticated public request", enabled: true, path: "/api/v1/settings/public", marker: "1"}, + {name: "unmarked shared API", enabled: true, path: "/api/v1/groups/available", role: "admin"}, + {name: "invalid marker", enabled: true, path: "/api/v1/groups/available", marker: "true", role: "admin"}, + {name: "admin prefix boundary", enabled: true, path: "/api/v1/administrator", role: "admin"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, tt.enabled, tt.path, tt.marker, tt.role, func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + header := recorder.Header().Get(servertiming.HeaderName) + if tt.wantHeader && header == "" { + t.Fatalf("%s header missing", servertiming.HeaderName) + } + if !tt.wantHeader && header != "" { + t.Fatalf("unexpected %s header: %q", servertiming.HeaderName, header) + } + if header != "" && (!strings.Contains(header, "total;dur=") || !strings.Contains(header, `cache;desc="bypass"`)) { + t.Fatalf("incomplete timing header: %q", header) + } + }) + } +} + +func TestServerTimingCollectorIsRequestScoped(t *testing.T) { + active := false + recorder := runServerTimingRequest(t, true, "/api/v1/keys", "1", "admin", func(c *gin.Context) { + active = servertiming.Active(c.Request.Context()) + c.Status(http.StatusNoContent) + }) + if !active { + t.Fatal("collector was not attached to marked request context") + } + if recorder.Header().Get(servertiming.HeaderName) == "" { + t.Fatal("timing header missing from status-only response") + } +} + +func TestServerTimingFinalizesBeforeEarlyCommit(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/stream", "", "admin", func(c *gin.Context) { + c.Status(http.StatusAccepted) + c.Writer.WriteHeaderNow() + }) + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatal("timing header was not written before response commit") + } +} + +func TestServerTimingFinalizesOnFlush(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/export", "", "admin", func(c *gin.Context) { + c.Writer.Flush() + }) + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatal("timing header was not written before stream flush") + } +} + +func TestServerTimingStatusResponses(t *testing.T) { + tests := []struct { + name string + status int + }{ + {name: "not modified", status: http.StatusNotModified}, + {name: "internal error", status: http.StatusInternalServerError}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/test", "", "admin", func(c *gin.Context) { + c.Status(tt.status) + }) + if recorder.Code != tt.status { + t.Fatalf("status = %d, want %d", recorder.Code, tt.status) + } + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatalf("timing header missing from status %d response", tt.status) + } + }) + } +} + +func TestServerTimingResponseWriterUnwraps(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + baseWriter := c.Writer + writer := &serverTimingResponseWriter{ResponseWriter: baseWriter} + if got := writer.Unwrap(); got != baseWriter { + t.Fatalf("Unwrap() = %T, want original Gin writer", got) + } +} + +func TestServerTimingCacheOutcome(t *testing.T) { + tests := []struct { + name string + headerName string + value string + want string + }{ + {name: "snapshot hit", headerName: snapshotCacheHeader, value: "hit", want: "hit"}, + {name: "usage miss", headerName: usageCacheHeader, value: "MISS", want: "miss"}, + {name: "invalid", headerName: snapshotCacheHeader, value: "stale", want: "bypass"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/dashboard", "", "admin", func(c *gin.Context) { + c.Header(tt.headerName, tt.value) + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + want := `cache;desc="` + tt.want + `"` + if got := recorder.Header().Get(servertiming.HeaderName); !strings.Contains(got, want) { + t.Fatalf("timing header %q does not contain %q", got, want) + } + }) + } +} + +func TestServerTimingResponseHeaderForWebSocket(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/admin/ops/ws/qps", nil) + collector := servertiming.New(time.Now()) + c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector)) + c.Set(string(ContextKeyUserRole), "admin") + + header := ServerTimingResponseHeader(c) + if header.Get(servertiming.HeaderName) == "" { + t.Fatal("WebSocket response header missing timing value") + } + + c.Set(string(ContextKeyUserRole), "user") + if got := ServerTimingResponseHeader(c); got != nil { + t.Fatalf("non-admin WebSocket received timing header: %#v", got) + } +} diff --git a/backend/internal/server/router.go b/backend/internal/server/router.go index 3d86373779..5fc70149fe 100644 --- a/backend/internal/server/router.go +++ b/backend/internal/server/router.go @@ -60,6 +60,7 @@ func SetupRouter( } return nil })) + r.Use(middleware2.ServerTiming(cfg.Server.EnableServerTiming)) // Serve embedded frontend with settings injection if available if web.HasEmbeddedFrontend() { diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 0d7e2a505a..5132022d4a 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -399,6 +399,7 @@ func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) { grok.POST("/oauth/exchange-code", h.Admin.GrokOAuth.ExchangeCode) grok.POST("/oauth/refresh-token", h.Admin.GrokOAuth.RefreshToken) grok.POST("/oauth/create-from-oauth", h.Admin.GrokOAuth.CreateAccountFromOAuth) + grok.POST("/sso-to-oauth", h.Admin.GrokOAuth.CreateAccountsFromSSO) grok.POST("/accounts/:id/refresh", h.Admin.GrokOAuth.RefreshAccountToken) grok.GET("/accounts/:id/quota", h.Admin.GrokOAuth.QueryQuota) grok.POST("/accounts/:id/reset-quota", h.Admin.GrokOAuth.ResetQuota) diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index ba5b4f61d1..45db227e58 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -84,6 +84,22 @@ func RegisterGatewayRoutes( }, }) } + videoEditHandler := func(c *gin.Context) { + if getGroupPlatform(c) == service.PlatformGrok { + h.OpenAIGateway.GrokVideoEdit(c) + return + } + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Videos API is not supported for this platform"}}) + } + videoExtensionHandler := func(c *gin.Context) { + if getGroupPlatform(c) == service.PlatformGrok { + h.OpenAIGateway.GrokVideoExtension(c) + return + } + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Videos API is not supported for this platform"}}) + } // API网关(Claude API兼容) gateway := r.Group("/v1") gateway.Use(bodyLimit) @@ -147,6 +163,7 @@ func RegisterGatewayRoutes( } h.Gateway.Responses(c) }) + gateway.POST("/alpha/search", h.OpenAIGateway.AlphaSearch) gateway.GET("/responses", func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) @@ -184,6 +201,8 @@ func RegisterGatewayRoutes( gateway.DELETE("/images/batches/:id", h.BatchImage.DeleteRecord) gateway.DELETE("/images/batches/:id/outputs", h.BatchImage.DeleteOutputs) gateway.POST("/videos/generations", videoGenerationHandler) + gateway.POST("/videos/edits", videoEditHandler) + gateway.POST("/videos/extensions", videoExtensionHandler) gateway.GET("/videos/:request_id", videoStatusHandler) } @@ -212,6 +231,7 @@ func RegisterGatewayRoutes( } r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) + r.POST("/alpha/search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch) r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) @@ -220,6 +240,7 @@ func RegisterGatewayRoutes( { codexDirect.POST("/responses", responsesHandler) codexDirect.POST("/responses/*subpath", responsesHandler) + codexDirect.POST("/alpha/search", h.OpenAIGateway.AlphaSearch) codexDirect.GET("/responses", func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) @@ -249,6 +270,8 @@ func RegisterGatewayRoutes( r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler) r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler) r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoGenerationHandler) + r.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoEditHandler) + r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoExtensionHandler) r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoStatusHandler) // Antigravity 模型列表 diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index 2779dd8f01..65c6824440 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -65,6 +65,36 @@ func TestGatewayRoutesOpenAIResponsesCompactPathIsRegistered(t *testing.T) { } } +func TestGatewayRoutesOpenAIAlphaSearchPathsAreRegistered(t *testing.T) { + router := newGatewayRoutesTestRouter() + registered := make(map[string]bool) + for _, route := range router.Routes() { + if route.Method == http.MethodPost { + registered[route.Path] = true + } + } + + for _, path := range []string{ + "/v1/alpha/search", + "/alpha/search", + "/backend-api/codex/alpha/search", + } { + require.True(t, registered[path], "POST %s should be registered", path) + } +} + +func TestGatewayRoutesAlphaSearchRejectsNonOpenAIGroup(t *testing.T) { + router := newGatewayRoutesTestRouter(service.PlatformGrok) + req := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusNotFound, w.Code) + require.Contains(t, w.Body.String(), "only available for OpenAI groups") +} + func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) { router := newGatewayRoutesTestRouter() @@ -93,6 +123,10 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) { "/images/edits", "/v1/videos/generations", "/videos/generations", + "/v1/videos/edits", + "/videos/edits", + "/v1/videos/extensions", + "/videos/extensions", } { req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"model":"grok-imagine","prompt":"draw a cat"}`)) req.Header.Set("Content-Type", "application/json") @@ -126,6 +160,10 @@ func TestGatewayRoutesNonGrokVideosAreRejectedAtPlatformGate(t *testing.T) { }{ {http.MethodPost, "/v1/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, {http.MethodPost, "/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, + {http.MethodPost, "/v1/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, + {http.MethodPost, "/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, + {http.MethodPost, "/v1/videos/extensions", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, + {http.MethodPost, "/videos/extensions", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, {http.MethodGet, "/v1/videos/request-123", ""}, {http.MethodGet, "/videos/request-123", ""}, } { diff --git a/backend/internal/server/routes/payment.go b/backend/internal/server/routes/payment.go index 7c54770028..9434a0b76d 100644 --- a/backend/internal/server/routes/payment.go +++ b/backend/internal/server/routes/payment.go @@ -28,7 +28,6 @@ func RegisterPaymentRoutes( authenticated.GET("/config", paymentHandler.GetPaymentConfig) authenticated.GET("/checkout-info", paymentHandler.GetCheckoutInfo) authenticated.GET("/plans", paymentHandler.GetPlans) - authenticated.GET("/channels", paymentHandler.GetChannels) authenticated.GET("/limits", paymentHandler.GetLimits) orders := authenticated.Group("/orders") diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index d099f93979..3e67fa5982 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -6,6 +6,7 @@ import ( "errors" "hash/fnv" "log/slog" + "net/url" "reflect" "sort" "strconv" @@ -82,6 +83,8 @@ type Account struct { type OpenAIEndpointCapability string +const openAILongContextBillingEnabledKey = "openai_long_context_billing_enabled" + const ( OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions" OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings" @@ -1191,6 +1194,14 @@ func (a *Account) IsOpenAI() bool { return a.Platform == PlatformOpenAI } +func (a *Account) IsOpenAILongContextBillingEnabled() bool { + if a == nil || !a.IsOpenAI() || a.Extra == nil { + return false + } + enabled, ok := a.Extra[openAILongContextBillingEnabledKey].(bool) + return ok && enabled +} + func (a *Account) IsAnthropic() bool { return a.Platform == PlatformAnthropic } @@ -1250,17 +1261,84 @@ func (a *Account) GetOpenAIRefreshToken() string { return a.GetCredential("refresh_token") } +// GetGrokBaseURL selects the upstream used by Grok text and Responses traffic. +// Grok media traffic has a different transport contract and must use +// GetGrokMediaBaseURL instead. func (a *Account) GetGrokBaseURL() string { if !a.IsGrok() { return "" } baseURL := a.GetCredential("base_url") + if a.IsGrokOAuth() { + if strings.TrimSpace(baseURL) == "" || isOfficialGrokAPIBaseURL(baseURL) { + return xai.DefaultCLIBaseURL + } + if _, err := xai.ValidateTrustedBaseURL(baseURL); err == nil { + return baseURL + } + return xai.DefaultCLIBaseURL + } if baseURL != "" { return baseURL } return xai.DefaultBaseURL } +// GetGrokMediaBaseURL selects the upstream used by Grok Imagine APIs. +// +// OAuth text requests need the CLI subscription proxy, but that proxy has a +// smaller request-body limit than the official Imagine API. Media requests can +// contain large base64 inputs, so default OAuth accounts must use api.x.ai. +// API-key accounts and explicit unsafe development overrides retain their +// configured base URL. +func (a *Account) GetGrokMediaBaseURL() string { + if !a.IsGrok() { + return "" + } + if !a.IsGrokOAuth() { + return a.GetGrokBaseURL() + } + + baseURL := a.GetCredential("base_url") + if strings.TrimSpace(baseURL) == "" || isOfficialGrokAPIBaseURL(baseURL) || isOfficialGrokCLIBaseURL(baseURL) { + return xai.DefaultBaseURL + } + if _, err := xai.ValidateTrustedBaseURL(baseURL); err == nil { + return baseURL + } + return xai.DefaultBaseURL +} + +func isOfficialGrokAPIBaseURL(raw string) bool { + return isOfficialGrokBaseURL(raw, xai.DefaultBaseURL) +} + +func isOfficialGrokCLIBaseURL(raw string) bool { + return isOfficialGrokBaseURL(raw, xai.DefaultCLIBaseURL) +} + +func isOfficialGrokBaseURL(raw, expected string) bool { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || parsed == nil || parsed.Opaque != "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return false + } + defaultURL, err := url.Parse(expected) + if err != nil { + return false + } + if !strings.EqualFold(parsed.Scheme, defaultURL.Scheme) || !strings.EqualFold(parsed.Hostname(), defaultURL.Hostname()) { + return false + } + if port := parsed.Port(); port != "" { + portNumber, err := strconv.Atoi(port) + if err != nil || portNumber != 443 { + return false + } + } + path := strings.TrimRight(parsed.Path, "/") + return path == "" || path == strings.TrimRight(defaultURL.Path, "/") +} + func (a *Account) GetGrokAccessToken() string { if !a.IsGrok() { return "" diff --git a/backend/internal/service/account_base_url_test.go b/backend/internal/service/account_base_url_test.go index a132219398..59f53db3db 100644 --- a/backend/internal/service/account_base_url_test.go +++ b/backend/internal/service/account_base_url_test.go @@ -4,6 +4,9 @@ package service import ( "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" ) func TestGetBaseURL(t *testing.T) { @@ -158,3 +161,249 @@ func TestGetGeminiBaseURL(t *testing.T) { }) } } + +func TestGetGrokBaseURLUsesSubscriptionProxyForOAuth(t *testing.T) { + tests := []struct { + name string + account Account + expected string + }{ + { + name: "oauth without base_url uses CLI subscription proxy", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{}, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API default is migrated at runtime to CLI subscription proxy", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": xai.DefaultBaseURL, + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API default with trailing slash is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": xai.DefaultBaseURL + "/", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API root is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://api.x.ai", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API root with canonical HTTPS port is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "HTTPS://API.X.AI:443/", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API canonical port with leading zeroes is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://api.x.ai:0443/v1", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API encoded version path is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://api.x.ai/%76%31", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth legacy API encoded trailing slash is migrated at runtime", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://api.x.ai/v1%2F", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "oauth non-default API port remains an explicit override", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://api.x.ai:8443/v1", + }, + }, + expected: "https://api.x.ai:8443/v1", + }, + { + name: "oauth explicit custom base_url stays pinned to CLI proxy by default", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://custom.example.com/v1", + }, + }, + expected: xai.DefaultCLIBaseURL, + }, + { + name: "API key without base_url uses official credit-backed API", + account: Account{ + Type: AccountTypeAPIKey, + Platform: PlatformGrok, + Credentials: map[string]any{}, + }, + expected: xai.DefaultBaseURL, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, tt.account.GetGrokBaseURL()) + }) + } +} + +func TestGetGrokBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnabled(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + account := Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://custom.example.com/v1", + }, + } + + require.Equal(t, "https://custom.example.com/v1", account.GetGrokBaseURL()) +} + +func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { + tests := []struct { + name string + account Account + expected string + }{ + { + name: "oauth without base_url uses official media API", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{}, + }, + expected: xai.DefaultBaseURL, + }, + { + name: "oauth stored CLI proxy uses official media API", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": xai.DefaultCLIBaseURL, + }, + }, + expected: xai.DefaultBaseURL, + }, + { + name: "oauth stored CLI proxy variant uses official media API", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "HTTPS://CLI-CHAT-PROXY.GROK.COM:443/%76%31/", + }, + }, + expected: xai.DefaultBaseURL, + }, + { + name: "oauth legacy official API remains on official media API", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": xai.DefaultBaseURL, + }, + }, + expected: xai.DefaultBaseURL, + }, + { + name: "oauth untrusted custom base_url is pinned to official media API", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://custom.example.com/v1", + }, + }, + expected: xai.DefaultBaseURL, + }, + { + name: "API key retains its configured media API", + account: Account{ + Type: AccountTypeAPIKey, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://grok.example.com/v1", + }, + }, + expected: "https://grok.example.com/v1", + }, + { + name: "non-Grok account has no Grok media base URL", + account: Account{ + Type: AccountTypeOAuth, + Platform: PlatformOpenAI, + Credentials: map[string]any{}, + }, + expected: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, tt.account.GetGrokMediaBaseURL()) + }) + } +} + +func TestGetGrokMediaBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnabled(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + account := Account{ + Type: AccountTypeOAuth, + Platform: PlatformGrok, + Credentials: map[string]any{ + "base_url": "https://custom.example.com/v1", + }, + } + + require.Equal(t, "https://custom.example.com/v1", account.GetGrokMediaBaseURL()) +} diff --git a/backend/internal/service/account_long_context_billing_test.go b/backend/internal/service/account_long_context_billing_test.go new file mode 100644 index 0000000000..709559d932 --- /dev/null +++ b/backend/internal/service/account_long_context_billing_test.go @@ -0,0 +1,290 @@ +//go:build unit + +package service + +import ( + "context" + "net/http" + "testing" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/stretchr/testify/require" +) + +func TestAccountIsOpenAILongContextBillingEnabled(t *testing.T) { + tests := []struct { + name string + account *Account + want bool + }{ + {name: "nil account is disabled", account: nil, want: false}, + {name: "non OpenAI account is disabled", account: &Account{Platform: PlatformGrok}, want: false}, + {name: "missing extra defaults disabled", account: &Account{Platform: PlatformOpenAI}, want: false}, + {name: "missing key defaults disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{}}, want: false}, + {name: "explicit true is enabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": true}}, want: true}, + {name: "explicit false is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": false}}, want: false}, + {name: "malformed value is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": "false"}}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, tt.account.IsOpenAILongContextBillingEnabled()) + }) + } +} + +func TestNormalizeOpenAILongContextBillingExtra(t *testing.T) { + t.Run("OpenAI missing key persists disabled default", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, nil) + + require.NoError(t, err) + require.Equal(t, false, extra["openai_long_context_billing_enabled"]) + }) + + t.Run("OpenAI explicit false is preserved", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": false}) + + require.NoError(t, err) + require.Equal(t, false, extra["openai_long_context_billing_enabled"]) + }) + + t.Run("OpenAI malformed value is rejected", func(t *testing.T) { + _, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": "false"}) + + require.Error(t, err) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + }) + + t.Run("non OpenAI extra is unchanged", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformGrok, nil) + + require.NoError(t, err) + require.Nil(t, extra) + }) + + t.Run("non OpenAI malformed value is ignored", func(t *testing.T) { + extra := map[string]any{openAILongContextBillingEnabledKey: "provider-owned"} + normalized, err := normalizeOpenAILongContextBillingExtra(PlatformAnthropic, extra) + + require.NoError(t, err) + require.Equal(t, extra, normalized) + }) +} + +type longContextBillingRepoStub struct { + accountRepoStub + account *Account + accounts []*Account + createdAccount *Account + updateExtraCalls int + bulkUpdateCalls int +} + +func (r *longContextBillingRepoStub) Create(_ context.Context, account *Account) error { + account.ID = 1 + r.account = account + r.createdAccount = account + return nil +} + +func (r *longContextBillingRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *longContextBillingRepoStub) GetByIDs(_ context.Context, _ []int64) ([]*Account, error) { + if r.accounts != nil { + return r.accounts, nil + } + if r.account == nil { + return nil, nil + } + return []*Account{r.account}, nil +} + +func (r *longContextBillingRepoStub) Update(_ context.Context, account *Account) error { + r.account = account + return nil +} + +func (r *longContextBillingRepoStub) UpdateExtra(_ context.Context, _ int64, _ map[string]any) error { + r.updateExtraCalls++ + return nil +} + +func (r *longContextBillingRepoStub) BulkUpdate(_ context.Context, _ []int64, _ AccountBulkUpdate) (int64, error) { + r.bulkUpdateCalls++ + return 1, nil +} + +func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingDisabled(t *testing.T) { + repo := &longContextBillingRepoStub{} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Name: "openai-account", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "test"}, + SkipDefaultGroupBind: true, + }) + + require.NoError(t, err) + require.Same(t, account, repo.createdAccount) + require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceCreateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Platform: PlatformOpenAI, + Extra: map[string]any{openAILongContextBillingEnabledKey: "false"}, + }) + + require.Nil(t, account) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Nil(t, repo.createdAccount) +} + +func TestAdminServiceUpdateAccountPreservesOpenAILongContextBillingOptOutWhenOmitted(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{openAILongContextBillingEnabledKey: false}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{}}) + + require.NoError(t, err) + require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountAllowsExplicitCodexImportOptIn(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "old-token"}, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: false, + "import_source": "codex_session", + }, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{ + Credentials: map[string]any{"access_token": "new-token"}, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: true, + "import_source": "codex_session", + }, + }) + + require.NoError(t, err) + require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountAllowsExplicitOptInOutsideCodexImport(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: false, + "import_source": "codex_session", + }, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{ + openAILongContextBillingEnabledKey: true, + "import_source": "codex_session", + }}) + + require.NoError(t, err) + require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{ + openAILongContextBillingEnabledKey: 1, + }}) + + require.Nil(t, account) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) +} + +func TestAdminServiceUpdateAccountExtraRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{ + openAILongContextBillingEnabledKey: "true", + }) + + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.updateExtraCalls) +} + +func TestAdminServiceUpdateAccountExtraAllowsProviderOwnedValueForNonOpenAIAccount(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformAnthropic}} + svc := &adminServiceImpl{accountRepo: repo} + + err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{ + openAILongContextBillingEnabledKey: "provider-owned", + }) + + require.NoError(t, err) + require.Equal(t, 1, repo.updateExtraCalls) +} + +func TestAdminServiceBulkUpdateAccountsRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: []bool{true}}, + }) + + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccountsAllowsProviderOwnedValueForNonOpenAIAccounts(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformGrok}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: []string{"provider-owned"}}, + }) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccountsRejectsMalformedValueForMixedTargetsIncludingOpenAI(t *testing.T) { + repo := &longContextBillingRepoStub{accounts: []*Account{ + {ID: 1, Platform: PlatformGrok}, + {ID: 2, Platform: PlatformOpenAI}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1, 2}, + Extra: map[string]any{openAILongContextBillingEnabledKey: "malformed"}, + }) + + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.bulkUpdateCalls) +} diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index bb57a93cd7..0e5b357dbb 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -587,8 +587,13 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account c.Writer.Header().Set("X-Accel-Buffering", "no") c.Writer.Flush() - // Create OpenAI Responses API payload - payload := createOpenAITestPayload(testModelID, isOAuth) + // Create OpenAI Responses API payload. OAuth accounts use ChatGPT Codex + // upstream and must apply the same model normalization as real forwarding. + upstreamTestModelID := testModelID + if isOAuth { + upstreamTestModelID = normalizeOpenAIModelForUpstream(credentialAccount, testModelID) + } + payload := createOpenAITestPayload(upstreamTestModelID, isOAuth) payloadBytes, _ := json.Marshal(payload) // Send test_start event once. A task-invalid Agent Identity response may @@ -683,31 +688,40 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account return s.processOpenAIStream(c, resp.Body) } -// testGrokAccountConnection tests a Grok OAuth account through xAI's Responses API. +// testGrokAccountConnection tests a Grok OAuth or API-key account through xAI's Responses API. func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *Account, modelID string) error { ctx := c.Request.Context() - if account.Type != AccountTypeOAuth { - return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported Grok account type: %s", account.Type)) - } - if s.grokTokenProvider == nil { - return s.sendErrorAndEnd(c, "Grok token provider not configured") - } if s.httpUpstream == nil { return s.sendErrorAndEnd(c, "HTTP upstream not configured") } testModelID := strings.TrimSpace(modelID) if testModelID == "" { - testModelID = "grok-4.3" + testModelID = grokDefaultResponsesModel } if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" { testModelID = mapped } - authToken, err := s.grokTokenProvider.GetAccessToken(ctx, account) - if err != nil { - return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to get Grok access token: %s", err.Error())) + var authToken string + switch account.Type { + case AccountTypeOAuth: + if s.grokTokenProvider == nil { + return s.sendErrorAndEnd(c, "Grok token provider not configured") + } + var err error + authToken, err = s.grokTokenProvider.GetAccessToken(ctx, account) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to get Grok access token: %s", err.Error())) + } + case AccountTypeAPIKey: + authToken = strings.TrimSpace(account.GetCredential("api_key")) + if authToken == "" { + return s.sendErrorAndEnd(c, "Grok API key is missing") + } + default: + return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported Grok account type: %s", account.Type)) } apiURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) @@ -741,7 +755,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json, text/event-stream") req.Header.Set("Authorization", "Bearer "+authToken) - req.Header.Set("User-Agent", "sub2api-grok/1.0") + applyGrokCLIHeaders(req.Header) proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -754,10 +768,19 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * } defer func() { _ = resp.Body.Close() }() - if snapshot := xai.ParseQuotaHeaders(resp.Header, resp.StatusCode); snapshot != nil && s.accountRepo != nil { + now := time.Now() + snapshot := parseGrokQuotaSnapshot(resp.Header, resp.StatusCode, now) + if snapshot != nil && s.accountRepo != nil { + resetAt, limited := grokRateLimitResetAt(snapshot, now) + if limited { + normalizeGrokExhaustedWindowResets(snapshot, resetAt, now) + } _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ grokQuotaSnapshotExtraKey: snapshot, }) + if limited { + persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + } } if resp.StatusCode != http.StatusOK { diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go index 356fc62fb7..4b0890ff44 100644 --- a/backend/internal/service/account_test_service_grok_test.go +++ b/backend/internal/service/account_test_service_grok_test.go @@ -3,6 +3,7 @@ package service import ( + "context" "io" "net/http" "net/http/httptest" @@ -15,6 +16,18 @@ import ( "github.com/tidwall/gjson" ) +type grokAccountTestRateLimitRepo struct { + *mockAccountRepoForGemini + rateLimitedCalls int + resetAt time.Time +} + +func (r *grokAccountTestRateLimitRepo) SetRateLimited(_ context.Context, _ int64, resetAt time.Time) error { + r.rateLimitedCalls++ + r.resetAt = resetAt + return nil +} + func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testing.T) { gin.SetMode(gin.TestMode) @@ -58,10 +71,122 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin err := svc.TestAccountConnection(c, account.ID, "grok", "", AccountTestModeDefault) require.NoError(t, err) - require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer grok-access-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) require.NotContains(t, rec.Body.String(), "claude") require.Contains(t, rec.Body.String(), `"model":"grok-4.3"`) require.Contains(t, rec.Body.String(), `"type":"test_complete"`) } + +func TestAccountTestService_TestAccountConnection_GrokDefaultsEmptyModelTo45(t *testing.T) { + gin.SetMode(gin.TestMode) + + account := &Account{ + ID: 16, + Name: "grok-oauth-default-model", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" + + "data: {\"type\":\"response.completed\"}\n\n", + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/16/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeDefault) + + require.NoError(t, err) + require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String()) + require.Contains(t, recorder.Body.String(), `"model":"grok-4.5"`) +} + +func TestAccountTestService_Grok429PersistsRateLimitReset(t *testing.T) { + gin.SetMode(gin.TestMode) + + account := &Account{ + ID: 14, + Name: "grok-oauth-limited", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + baseRepo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + repo := &grokAccountTestRateLimitRepo{mockAccountRepoForGemini: baseRepo} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"45"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/14/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "grok", "", AccountTestModeDefault) + + require.Error(t, err) + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, time.Now().Add(45*time.Second), repo.resetAt, time.Second) +} + +func TestAccountTestService_Grok429WithoutQuotaHeadersUsesFallback(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 15, Name: "grok-oauth-limited-no-headers", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + baseRepo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + repo := &grokAccountTestRateLimitRepo{mockAccountRepoForGemini: baseRepo} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"quota exhausted"}}`)), + }} + svc := &AccountTestService{ + accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, nil), httpUpstream: upstream, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/15/test", nil) + before := time.Now() + + err := svc.TestAccountConnection(c, account.ID, "grok", "", AccountTestModeDefault) + + require.Error(t, err) + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, before.Add(grokRateLimitFallbackCooldown), repo.resetAt, time.Second) +} diff --git a/backend/internal/service/account_test_service_openai_test.go b/backend/internal/service/account_test_service_openai_test.go index af28085123..083d882ea7 100644 --- a/backend/internal/service/account_test_service_openai_test.go +++ b/backend/internal/service/account_test_service_openai_test.go @@ -137,6 +137,34 @@ func TestAccountTestService_OpenAISuccessPersistsSnapshotFromHeaders(t *testing. require.Contains(t, recorder.Body.String(), "test_complete") } +func TestAccountTestService_OpenAIOAuthTestNormalizesGPT56Alias(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := newTestContext() + + resp := newJSONResponse(http.StatusOK, "") + resp.Body = io.NopCloser(strings.NewReader(`data: {"type":"response.completed"} + +`)) + + upstream := &queuedHTTPUpstream{responses: []*http.Response{resp}} + svc := &AccountTestService{httpUpstream: upstream} + account := &Account{ + ID: 90, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{"access_token": "test-token"}, + } + + err := svc.testOpenAIAccountConnection(ctx, account, "gpt-5.6", "", "") + require.NoError(t, err) + require.Len(t, upstream.requests, 1) + + body, err := io.ReadAll(upstream.requests[0].Body) + require.NoError(t, err) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(body, "model").String()) +} + func TestAccountTestService_OpenAIShadowUsesParentCredentialsAndShadowModel(t *testing.T) { gin.SetMode(gin.TestMode) ctx, recorder := newTestContext() diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 622323c617..12dc25a673 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -111,6 +111,8 @@ const ( apiQueryMaxJitter = 800 * time.Millisecond // 用量查询最大随机延迟 windowStatsCacheTTL = 1 * time.Minute openAIProbeCacheTTL = 10 * time.Minute + grokProbeRetryTTL = 1 * time.Minute + grokFreeQuotaWindow = 24 * time.Hour openAICodexProbeVersion = "0.144.1" ) @@ -122,6 +124,7 @@ type UsageCache struct { apiFlight singleflight.Group // 防止同一账号的并发请求击穿缓存(Anthropic) antigravityFlight singleflight.Group // 防止同一 Antigravity 账号的并发请求击穿缓存 openAIProbeCache sync.Map // accountID -> time.Time + grokProbeCache sync.Map // accountID -> last billing probe attempt } // NewUsageCache 创建 UsageCache 实例 @@ -196,15 +199,19 @@ type UsageInfo struct { AntigravityQuota map[string]*AntigravityModelQuota `json:"antigravity_quota,omitempty"` // Grok / xAI 被动额度快照 - GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` - GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` - GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` - GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` - GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` - GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` - GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` - GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` - GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` + GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` + GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` + GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` + GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` + GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` + GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` + GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` + GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + GrokLocalUsage24h *WindowStats `json:"grok_local_usage_24h,omitempty"` + GrokLocalUsage7d *WindowStats `json:"grok_local_usage_7d,omitempty"` + GrokLocalUsageMonthly *WindowStats `json:"grok_local_usage_monthly,omitempty"` + GrokBilling *xai.BillingSummary `json:"grok_billing,omitempty"` // Antigravity 账号级信息 SubscriptionTier string `json:"subscription_tier,omitempty"` // 归一化订阅等级: FREE/PRO/ULTRA/UNKNOWN @@ -287,6 +294,7 @@ type AccountUsageService struct { geminiQuotaService *GeminiQuotaService antigravityQuotaFetcher *AntigravityQuotaFetcher grokQuotaFetcher *GrokQuotaFetcher + grokQuotaService *GrokQuotaService openAIQuotaService *OpenAIQuotaService cache *UsageCache identityCache IdentityCache @@ -303,6 +311,7 @@ func NewAccountUsageService( geminiQuotaService *GeminiQuotaService, antigravityQuotaFetcher *AntigravityQuotaFetcher, grokQuotaFetcher *GrokQuotaFetcher, + grokQuotaService *GrokQuotaService, openAIQuotaService *OpenAIQuotaService, cache *UsageCache, identityCache IdentityCache, @@ -315,6 +324,7 @@ func NewAccountUsageService( geminiQuotaService: geminiQuotaService, antigravityQuotaFetcher: antigravityQuotaFetcher, grokQuotaFetcher: grokQuotaFetcher, + grokQuotaService: grokQuotaService, openAIQuotaService: openAIQuotaService, cache: cache, identityCache: identityCache, @@ -360,8 +370,8 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for } if account.Platform == PlatformGrok { - usage, err := s.getGrokUsage(ctx, account) - if err == nil { + usage, err := s.getGrokUsage(ctx, account, forceProbe) + if err == nil && usage != nil && usage.Error == "" { s.tryClearRecoverableAccountError(ctx, account) } return usage, err @@ -947,11 +957,21 @@ func (s *AccountUsageService) getAntigravityUsage(ctx context.Context, account * return usage, nil } -func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account) (*UsageInfo, error) { +func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account, force bool) (*UsageInfo, error) { if s.grokQuotaFetcher == nil { now := time.Now() return &UsageInfo{UpdatedAt: &now}, nil } + var billingProbeResult *GrokQuotaProbeResult + if account != nil && account.IsGrokOAuth() && s.grokQuotaService != nil && (force || grokBillingSnapshotNeedsRefresh(account, time.Now())) && s.shouldProbeGrokBilling(account.ID, time.Now(), force) { + result, err := s.grokQuotaService.ProbeBilling(ctx, account.ID) + if err == nil && result != nil && result.Billing != nil { + billingProbeResult = result + mergeAccountExtra(account, map[string]any{grokBillingExtraKey: result.Billing}) + } else if err != nil && force { + return nil, err + } + } usage := s.grokQuotaFetcher.BuildUsageInfo(account) if usage.GrokQuotaSnapshotState == "" { if usage.ErrorCode == "quota_unknown" { @@ -961,9 +981,20 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account } } - if s.usageLogRepo != nil && account != nil { - if stats, err := s.usageLogRepo.GetAccountTodayStats(ctx, account.ID); err == nil && stats != nil { - usage.GrokLocalUsage = windowStatsFromAccountStats(stats) + if account != nil { + if s.usageLogRepo != nil { + if stats, err := s.usageLogRepo.GetAccountTodayStats(ctx, account.ID); err == nil && stats != nil { + usage.GrokLocalUsage = windowStatsFromAccountStats(stats) + } + } + if billingProbeResult != nil { + usage.GrokLocalUsage24h = billingProbeResult.LocalUsage24h + usage.GrokLocalUsage7d = billingProbeResult.LocalUsage7d + usage.GrokLocalUsageMonthly = billingProbeResult.LocalUsageMonthly + } else if s.usageLogRepo != nil { + usage.GrokLocalUsage24h, usage.GrokLocalUsage7d, usage.GrokLocalUsageMonthly = grokLocalUsageForQuota( + ctx, s.usageLogRepo, account.ID, usage.GrokBilling, time.Now().UTC(), + ) } } @@ -971,6 +1002,110 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account return usage, nil } +func grokLocalUsageForQuota( + ctx context.Context, + repo UsageLogRepository, + accountID int64, + billing *xai.BillingSummary, + now time.Time, +) (*WindowStats, *WindowStats, *WindowStats) { + if grokBillingHasAuthoritativeQuota(billing) { + weekly, monthly := grokLocalUsageForBilling(ctx, repo, accountID, billing, now) + return nil, weekly, monthly + } + return grokLocalUsage24h(ctx, repo, accountID, now), nil, nil +} + +func grokLocalUsage24h(ctx context.Context, repo UsageLogRepository, accountID int64, now time.Time) *WindowStats { + if repo == nil || accountID <= 0 { + return nil + } + start := now.UTC().Add(-grokFreeQuotaWindow) + stats, err := repo.GetAccountWindowStats(ctx, accountID, start) + if err != nil { + slog.Warn("grok_rolling_24h_usage_query_failed", "account_id", accountID, "window_start", start, "error", err) + return nil + } + return windowStatsFromAccountStats(stats) +} + +func grokLocalUsageForBilling( + ctx context.Context, + repo UsageLogRepository, + accountID int64, + billing *xai.BillingSummary, + now time.Time, +) (*WindowStats, *WindowStats) { + var weekly *WindowStats + var monthly *WindowStats + if repo == nil || accountID <= 0 { + return weekly, monthly + } + if start, ok := currentGrokBillingWindow(billing, true, now); ok { + if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil { + weekly = windowStatsFromAccountStats(stats) + } else { + slog.Warn("grok_window_usage_query_failed", "account_id", accountID, "window_start", start, "error", err) + } + } + if start, ok := currentGrokBillingWindow(billing, false, now); ok { + if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil { + monthly = windowStatsFromAccountStats(stats) + } else { + slog.Warn("grok_monthly_usage_query_failed", "account_id", accountID, "window_start", start, "error", err) + } + } + return weekly, monthly +} + +func currentGrokBillingWindow(billing *xai.BillingSummary, weekly bool, now time.Time) (time.Time, bool) { + if billing == nil { + return time.Time{}, false + } + startRaw, endRaw := billing.BillingPeriodStart, billing.BillingPeriodEnd + if weekly { + if billing.PeriodType != "weekly" { + return time.Time{}, false + } + startRaw, endRaw = billing.PeriodStart, billing.PeriodEnd + } + start, startErr := parseTime(strings.TrimSpace(startRaw)) + end, endErr := parseTime(strings.TrimSpace(endRaw)) + if startErr != nil || endErr != nil || now.Before(start) || !now.Before(end) { + return time.Time{}, false + } + return start, true +} + +func grokBillingSnapshotNeedsRefresh(account *Account, now time.Time) bool { + if account == nil { + return false + } + billing, err := grokBillingSnapshotFromExtra(account.Extra) + if err != nil || billing == nil || billing.Partial || len(billing.FailedWindows) > 0 { + return true + } + stamp := strings.TrimSpace(billing.UpdatedAt) + if stamp == "" { + stamp = strings.TrimSpace(billing.FetchedAt) + } + updatedAt, err := parseTime(stamp) + return err != nil || now.Sub(updatedAt) >= openAIProbeCacheTTL +} + +func (s *AccountUsageService) shouldProbeGrokBilling(accountID int64, now time.Time, force bool) bool { + if force || s == nil || s.cache == nil || accountID <= 0 { + return true + } + if cached, ok := s.cache.grokProbeCache.Load(accountID); ok { + if ts, ok := cached.(time.Time); ok && now.Sub(ts) < grokProbeRetryTTL { + return false + } + } + s.cache.grokProbeCache.Store(accountID, now) + return true +} + // recalcAntigravityRemainingSeconds 重新计算 Antigravity UsageInfo 中各窗口的 RemainingSeconds // 用于从缓存取出时更新倒计时,避免返回过时的剩余秒数 func recalcAntigravityRemainingSeconds(info *UsageInfo) { diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 52e5ce719b..8cb6d8e63b 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "log/slog" + "maps" "net/http" "strconv" "strings" @@ -68,7 +69,65 @@ func normalizeAccountConcurrency(platform, accountType string, concurrency int) return concurrency } +// ValidateOpenAILongContextBillingExtra validates the OpenAI account billing flag when present. +func ValidateOpenAILongContextBillingExtra(platform string, extra map[string]any) error { + if platform != PlatformOpenAI { + return nil + } + raw, exists := extra[openAILongContextBillingEnabledKey] + if !exists { + return nil + } + if _, ok := raw.(bool); !ok { + return infraerrors.BadRequest( + "OPENAI_LONG_CONTEXT_BILLING_INVALID", + "openai_long_context_billing_enabled must be a boolean", + ) + } + return nil +} + +func normalizeOpenAILongContextBillingExtra(platform string, extra map[string]any) (map[string]any, error) { + if platform != PlatformOpenAI { + return extra, nil + } + if err := ValidateOpenAILongContextBillingExtra(platform, extra); err != nil { + return nil, err + } + + normalized := maps.Clone(extra) + if normalized == nil { + normalized = make(map[string]any, 1) + } + _, exists := normalized[openAILongContextBillingEnabledKey] + if !exists { + normalized[openAILongContextBillingEnabledKey] = false + } + return normalized, nil +} + +func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *UpdateAccountInput) (map[string]any, error) { + normalized, err := normalizeOpenAILongContextBillingExtra(account.Platform, input.Extra) + if err != nil || account.Platform != PlatformOpenAI { + return normalized, err + } + + _, provided := input.Extra[openAILongContextBillingEnabledKey] + current, hasCurrent := account.Extra[openAILongContextBillingEnabledKey].(bool) + if !provided { + if hasCurrent { + normalized[openAILongContextBillingEnabledKey] = current + } + } + return normalized, nil +} + func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) { + accountExtra, err := normalizeOpenAILongContextBillingExtra(input.Platform, input.Extra) + if err != nil { + return nil, err + } + // 绑定分组 groupIDs := input.GroupIDs // 如果没有指定分组,自动绑定对应平台的默认分组 @@ -103,7 +162,7 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou Platform: input.Platform, Type: input.Type, Credentials: input.Credentials, - Extra: input.Extra, + Extra: accountExtra, ProxyID: input.ProxyID, Concurrency: normalizeAccountConcurrency(input.Platform, input.Type, input.Concurrency), Priority: input.Priority, @@ -183,6 +242,13 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if err != nil { return nil, err } + var normalizedExtra map[string]any + if input.Extra != nil { + normalizedExtra, err = normalizeOpenAILongContextBillingUpdateExtra(account, input) + if err != nil { + return nil, err + } + } // 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用, // 必须在此守住,否则仅在创建时的保证可被这些路径绕过。 if account.IsCredentialShadow() { @@ -238,10 +304,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // 保留配额用量字段,防止编辑账号时意外重置 for _, key := range []string{"quota_used", "quota_daily_used", "quota_daily_start", "quota_weekly_used", "quota_weekly_start"} { if v, ok := account.Extra[key]; ok { - input.Extra[key] = v + normalizedExtra[key] = v } } - account.Extra = input.Extra + account.Extra = normalizedExtra if account.Platform == PlatformAntigravity && wasOveragesEnabled && !account.IsOveragesEnabled() { delete(account.Extra, "antigravity_credits_overages") // 清理旧版 overages 运行态 // 清除 AICredits 限流 key @@ -353,6 +419,15 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // UpdateAccountExtra 仅对 Extra JSONB 做 key 级合并,避免覆盖其它运行态键 // (如 model_rate_limits / passive_usage_* 等)。 func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error { + if _, exists := updates[openAILongContextBillingEnabledKey]; exists { + account, err := s.accountRepo.GetByID(ctx, id) + if err != nil { + return err + } + if err := ValidateOpenAILongContextBillingExtra(account.Platform, updates); err != nil { + return err + } + } if len(updates) == 0 { return nil } @@ -386,16 +461,28 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp } needMixedChannelCheck := input.GroupIDs != nil && !input.SkipMixedChannelCheck + _, hasLongContextBillingUpdate := input.Extra[openAILongContextBillingEnabledKey] // 预取所有目标账号,供凭据守卫/代理守卫/混合渠道检查共用,避免多次 DB 查询。 var cachedTargets []*Account - if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck { + if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate { loaded, err := s.accountRepo.GetByIDs(ctx, input.AccountIDs) if err != nil { return nil, err } cachedTargets = loaded } + if hasLongContextBillingUpdate { + for _, account := range cachedTargets { + if account == nil || account.Platform != PlatformOpenAI { + continue + } + if err := ValidateOpenAILongContextBillingExtra(account.Platform, input.Extra); err != nil { + return nil, err + } + break + } + } // 影子账号绝不持有凭据:批量更新携带凭据时,目标中不得含影子(外审 G5,与单账号 // UpdateAccount 守卫对齐)。覆盖显式 IDs 与 filter 解析出的 IDs(此处 AccountIDs 已解析完成)。 @@ -745,6 +832,9 @@ func (s *adminServiceImpl) CreateShadow(ctx context.Context, parentID int64, opt Priority: priority, Concurrency: concurrency, Schedulable: true, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: parent.IsOpenAILongContextBillingEnabled(), + }, } // 5. 持久化(Create 填充 shadow.ID)。并发竞态:预查(步骤2)放行后另一请求抢先建成,本次会撞 diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index c85056d623..d622d547b0 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -156,6 +156,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn videoPrice480P := normalizePrice(input.VideoPrice480P) videoPrice720P := normalizePrice(input.VideoPrice720P) videoPrice1080P := normalizePrice(input.VideoPrice1080P) + webSearchPricePerCall := normalizePrice(input.WebSearchPricePerCall) imageRateMultiplier := 1.0 if input.ImageRateMultiplier != nil { if *input.ImageRateMultiplier < 0 { @@ -287,6 +288,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn VideoPrice480P: videoPrice480P, VideoPrice720P: videoPrice720P, VideoPrice1080P: videoPrice1080P, + WebSearchPricePerCall: webSearchPricePerCall, ClaudeCodeOnly: input.ClaudeCodeOnly, FallbackGroupID: input.FallbackGroupID, FallbackGroupIDOnInvalidRequest: fallbackOnInvalidRequest, @@ -543,6 +545,9 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.VideoPrice1080P != nil { group.VideoPrice1080P = normalizePrice(input.VideoPrice1080P) } + if input.WebSearchPricePerCall != nil { + group.WebSearchPricePerCall = normalizePrice(input.WebSearchPricePerCall) + } // Claude Code 客户端限制 if input.ClaudeCodeOnly != nil { diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 07b85ab827..2a7125f51f 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -220,8 +220,10 @@ type CreateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 - ClaudeCodeOnly bool // 仅允许 Claude Code 客户端 - FallbackGroupID *int64 // 降级分组 ID + // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用);nil/负数按默认价 0.01 处理 + WebSearchPricePerCall *float64 + ClaudeCodeOnly bool // 仅允许 Claude Code 客户端 + FallbackGroupID *int64 // 降级分组 ID // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) @@ -274,8 +276,10 @@ type UpdateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 - ClaudeCodeOnly *bool // 仅允许 Claude Code 客户端 - FallbackGroupID *int64 // 降级分组 ID + // Codex alpha/search 网页搜索单次价格(USD/次);nil 表示不修改,负数表示清除回默认价 0.01 + WebSearchPricePerCall *float64 + ClaudeCodeOnly *bool // 仅允许 Claude Code 客户端 + FallbackGroupID *int64 // 降级分组 ID // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) diff --git a/backend/internal/service/admin_service_spark_shadow_test.go b/backend/internal/service/admin_service_spark_shadow_test.go index 6b4017207a..0eda0d93c7 100644 --- a/backend/internal/service/admin_service_spark_shadow_test.go +++ b/backend/internal/service/admin_service_spark_shadow_test.go @@ -157,6 +157,39 @@ func TestCreateShadow(t *testing.T) { require.Error(t, err) } +func TestCreateShadowInheritsParentEffectiveOpenAILongContextBillingValue(t *testing.T) { + tests := []struct { + name string + parentExtra map[string]any + want bool + }{ + {name: "missing parent value defaults disabled", want: false}, + {name: "explicit parent opt-out is inherited", parentExtra: map[string]any{openAILongContextBillingEnabledKey: false}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + repo := newSparkShadowRepoStub() + svc := &adminServiceImpl{accountRepo: repo} + parent := &Account{ + Name: "parent", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Credentials: map[string]any{"access_token": "token"}, + Extra: tt.parentExtra, + } + require.NoError(t, repo.Create(context.Background(), parent)) + + shadow, err := svc.CreateShadow(context.Background(), parent.ID, ShadowOptions{Name: "shadow"}) + + require.NoError(t, err) + require.Equal(t, tt.want, shadow.Extra[openAILongContextBillingEnabledKey]) + require.Equal(t, tt.want, shadow.IsOpenAILongContextBillingEnabled()) + }) + } +} + // TestCreateShadow_BindGroups は BindGroups の後置呼び出しを検証する。 // 影子账号が指定グループに属し、ListSchedulableByGroupID で取得可能であること。 func TestCreateShadow_BindGroups(t *testing.T) { diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index 11b5246a1d..0cd1c6d584 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -78,6 +78,7 @@ type APIKeyAuthGroupSnapshot struct { VideoPrice480P *float64 `json:"video_price_480p,omitempty"` VideoPrice720P *float64 `json:"video_price_720p,omitempty"` VideoPrice1080P *float64 `json:"video_price_1080p,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call,omitempty"` ClaudeCodeOnly bool `json:"claude_code_only"` FallbackGroupID *int64 `json:"fallback_group_id,omitempty"` FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request,omitempty"` diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index 539c7375d9..5c45408659 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -14,7 +14,7 @@ import ( "github.com/dgraph-io/ristretto" ) -const apiKeyAuthSnapshotVersion = 14 // v14: include group video pricing fields +const apiKeyAuthSnapshotVersion = 15 // v15: include group web search per-call pricing type apiKeyAuthCacheConfig struct { l1Size int @@ -270,6 +270,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey) VideoPrice480P: apiKey.Group.VideoPrice480P, VideoPrice720P: apiKey.Group.VideoPrice720P, VideoPrice1080P: apiKey.Group.VideoPrice1080P, + WebSearchPricePerCall: apiKey.Group.WebSearchPricePerCall, ClaudeCodeOnly: apiKey.Group.ClaudeCodeOnly, FallbackGroupID: apiKey.Group.FallbackGroupID, FallbackGroupIDOnInvalidRequest: apiKey.Group.FallbackGroupIDOnInvalidRequest, @@ -353,6 +354,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho VideoPrice480P: snapshot.Group.VideoPrice480P, VideoPrice720P: snapshot.Group.VideoPrice720P, VideoPrice1080P: snapshot.Group.VideoPrice1080P, + WebSearchPricePerCall: snapshot.Group.WebSearchPricePerCall, ClaudeCodeOnly: snapshot.Group.ClaudeCodeOnly, FallbackGroupID: snapshot.Group.FallbackGroupID, FallbackGroupIDOnInvalidRequest: snapshot.Group.FallbackGroupIDOnInvalidRequest, diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 3dfd500b05..7fa69d41ea 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -153,14 +153,15 @@ type UsageTokens struct { // CostBreakdown 费用明细 type CostBreakdown struct { - InputCost float64 - OutputCost float64 - ImageOutputCost float64 - CacheCreationCost float64 - CacheReadCost float64 - TotalCost float64 - ActualCost float64 // 应用倍率后的实际费用 - BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充 + InputCost float64 + OutputCost float64 + ImageOutputCost float64 + CacheCreationCost float64 + CacheReadCost float64 + TotalCost float64 + ActualCost float64 // 应用倍率后的实际费用 + BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充 + LongContextBillingApplied bool } // ErrModelPricingUnavailable indicates that none of the configured pricing @@ -563,15 +564,19 @@ func (s *BillingService) initFallbackPricing() { s.fallbackPrices["grok-4.3"] = &ModelPricing{ InputPricePerToken: 1.25e-6, OutputPricePerToken: 2.5e-6, - CacheReadPricePerToken: 0, + CacheReadPricePerToken: 0.2e-6, SupportsCacheBreakdown: false, LongContextInputThreshold: 1000000, LongContextInputMultiplier: 1, } - // xAI Grok Build 0.1 (official docs: $1 input / $2 output per MTok) + // xAI Grok Build 0.1 (official docs: $1 input / $0.20 cached input / + // $2 output per MTok). Composer is available only through Grok Build and + // has no standalone public API rate card, so its aliases use this coding + // model rate instead of silently billing at zero. s.fallbackPrices["grok-build-0.1"] = &ModelPricing{ InputPricePerToken: 1e-6, OutputPricePerToken: 2e-6, + CacheReadPricePerToken: 0.2e-6, SupportsCacheBreakdown: false, } } @@ -745,9 +750,14 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { switch modelLower { case "grok", "grok-latest", "grok-4.5", "grok-4.5-latest", "grok-build-latest": return s.fallbackPrices["grok-4.5"] - case "grok-4.3": + case "grok-4.3", + "grok-4.20-0309-reasoning", + "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent-0309", + "grok-4.20-reasoning", + "grok-4.20-non-reasoning": return s.fallbackPrices["grok-4.3"] - case "grok-build", "grok-build-0.1": + case "grok-build", "grok-build-0.1", "grok-composer", "grok-composer-2.5-fast", "composer-2.5": return s.fallbackPrices["grok-build-0.1"] } @@ -856,16 +866,17 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing // CostInput 统一计费输入 type CostInput struct { - Ctx context.Context - Model string - GroupID *int64 // 用于渠道定价查找 - Tokens UsageTokens - RequestCount int // 按次计费时使用 - SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) - RateMultiplier float64 - ServiceTier string // "priority","flex","" 等 - Resolver *ModelPricingResolver // 定价解析器 - Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) + Ctx context.Context + Model string + GroupID *int64 // 用于渠道定价查找 + Tokens UsageTokens + RequestCount int // 按次计费时使用 + SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) + RateMultiplier float64 + ServiceTier string // "priority","flex","" 等 + Resolver *ModelPricingResolver // 定价解析器 + Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) + LongContextBillingEnabled *bool } // CalculateCostUnified 统一计费入口,支持三种计费模式。 @@ -873,7 +884,18 @@ type CostInput struct { func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, error) { if input.Resolver == nil { // 无 Resolver,回退到旧路径 - return s.calculateCostInternal(input.Model, input.Tokens, input.RateMultiplier, input.ServiceTier, nil) + applyLongContextBilling := true + if input.LongContextBillingEnabled != nil { + applyLongContextBilling = *input.LongContextBillingEnabled + } + return s.calculateCostInternalWithPolicy( + input.Model, + input.Tokens, + input.RateMultiplier, + input.ServiceTier, + nil, + applyLongContextBilling, + ) } // 优先使用预解析结果,避免重复 Resolve 调用 @@ -920,6 +942,9 @@ func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input Cos // 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层) applyLongCtx := len(resolved.Intervals) == 0 + if input.LongContextBillingEnabled != nil { + applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled + } return s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx), nil } @@ -960,7 +985,10 @@ func (s *BillingService) computeTokenBreakdown( tierMultiplier = serviceTierCostMultiplier(serviceTier) } - if applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) { + longContextPricingEligible := applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) + var baselineCost *CostBreakdown + if longContextPricingEligible { + baselineCost = s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, false) inputPrice *= pricing.LongContextInputMultiplier outputPrice *= pricing.LongContextOutputMultiplier // 缓存读取本质上是输入侧的复用,应与 input 一同应用长上下文倍率; @@ -1024,6 +1052,7 @@ func (s *BillingService) computeTokenBreakdown( bd.TotalCost = bd.InputCost + bd.OutputCost + bd.ImageOutputCost + bd.CacheCreationCost + bd.CacheReadCost bd.ActualCost = bd.TotalCost * rateMultiplier + bd.LongContextBillingApplied = baselineCost != nil && bd.ActualCost > baselineCost.ActualCost return bd } @@ -1083,7 +1112,28 @@ func (s *BillingService) CalculateCostWithServiceTier(model string, tokens Usage return s.calculateCostInternal(model, tokens, rateMultiplier, serviceTier, nil) } +func (s *BillingService) calculateCostWithServiceTierPolicy( + model string, + tokens UsageTokens, + rateMultiplier float64, + serviceTier string, + longContextBillingEnabled bool, +) (*CostBreakdown, error) { + return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, nil, longContextBillingEnabled) +} + func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens, rateMultiplier float64, serviceTier string, channelPricing *ChannelModelPricing) (*CostBreakdown, error) { + return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, channelPricing, true) +} + +func (s *BillingService) calculateCostInternalWithPolicy( + model string, + tokens UsageTokens, + rateMultiplier float64, + serviceTier string, + channelPricing *ChannelModelPricing, + longContextBillingEnabled bool, +) (*CostBreakdown, error) { var pricing *ModelPricing var err error if channelPricing != nil { @@ -1095,8 +1145,7 @@ func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens, return nil, err } - // 旧路径始终检查长上下文定价(无区间定价概念) - return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, true), nil + return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, longContextBillingEnabled), nil } func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing *ModelPricing) *ModelPricing { @@ -1227,13 +1276,14 @@ func (s *BillingService) CalculateCostWithLongContext(model string, tokens Usage // 合并成本 return &CostBreakdown{ - InputCost: inRangeCost.InputCost + outRangeCost.InputCost, - OutputCost: inRangeCost.OutputCost, - ImageOutputCost: inRangeCost.ImageOutputCost, - CacheCreationCost: inRangeCost.CacheCreationCost, - CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost, - TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost, - ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost, + InputCost: inRangeCost.InputCost + outRangeCost.InputCost, + OutputCost: inRangeCost.OutputCost, + ImageOutputCost: inRangeCost.ImageOutputCost, + CacheCreationCost: inRangeCost.CacheCreationCost, + CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost, + TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost, + ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost, + LongContextBillingApplied: outRangeCost.ActualCost > 0, }, nil } @@ -1320,8 +1370,36 @@ const ( defaultGrokImagineVideo15Price480P = 0.08 defaultGrokImagineVideo15Price720P = 0.14 defaultGrokImagineVideo15Price1080P = 0.25 + + // Codex alpha/search 网页搜索单次默认价:OpenAI 官方 web search 定价 $10/1000 次。 + defaultWebSearchPricePerCall = 0.01 ) +// CalculateWebSearchCost 计算 Codex alpha/search 网页搜索按次费用。 +// callCount: 搜索调用次数(每次请求为 1) +// groupPrice: 分组配置的单次价格(nil 表示使用默认价 0.01;0 表示免费) +// rateMultiplier: 分组费率倍数 +func (s *BillingService) CalculateWebSearchCost(callCount int, groupPrice *float64, rateMultiplier float64) *CostBreakdown { + if callCount <= 0 { + return &CostBreakdown{} + } + unitPrice := defaultWebSearchPricePerCall + if groupPrice != nil && *groupPrice >= 0 { + unitPrice = *groupPrice + } + totalCost := unitPrice * float64(callCount) + + // 应用倍率(保存时强制 > 0;负数按 0 处理避免按 1x 误扣) + if rateMultiplier < 0 { + rateMultiplier = 0 + } + return &CostBreakdown{ + TotalCost: totalCost, + ActualCost: totalCost * rateMultiplier, + BillingMode: string(BillingModePerRequest), + } +} + // CalculateImageCost 计算图片生成费用 // model: 请求的模型名称(用于获取 LiteLLM 默认价格) // imageSize: 图片尺寸 "1K", "2K", "4K" diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index c1f3f6e557..885da194e3 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -261,6 +261,23 @@ func TestCalculateCost_OpenAIGPT54LongContextAppliesWholeSessionMultipliers(t *t require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10) + require.True(t, cost.LongContextBillingApplied) +} + +func TestCalculateCost_OpenAIGPT54LongContextMarkerRequiresActualCostIncrease(t *testing.T) { + svc := newTestBillingService() + + cost, err := svc.calculateCostWithServiceTierPolicy( + "gpt-5.4-2026-03-05", + UsageTokens{InputTokens: 300000}, + 0, + "", + true, + ) + + require.NoError(t, err) + require.Zero(t, cost.ActualCost) + require.False(t, cost.LongContextBillingApplied) } func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) { @@ -831,6 +848,17 @@ func TestCalculateCostWithLongContext_AboveThreshold_CacheBelowThreshold(t *test require.True(t, cost.ActualCost > normalCost.ActualCost, "长上下文费用应高于正常费用") } +func TestCalculateCostWithLongContext_MarkerRequiresActualCostIncrease(t *testing.T) { + svc := newTestBillingService() + tokens := UsageTokens{InputTokens: 300000} + + cost, err := svc.CalculateCostWithLongContext("claude-sonnet-4", tokens, 0, 200000, 2.0) + + require.NoError(t, err) + require.Zero(t, cost.ActualCost) + require.False(t, cost.LongContextBillingApplied) +} + func TestCalculateCostWithLongContext_DisabledThreshold(t *testing.T) { svc := newTestBillingService() @@ -1039,6 +1067,58 @@ func TestGetModelPricing_Grok45OfficialFallback(t *testing.T) { } } +func TestGetModelPricing_GrokCatalogFallbacks(t *testing.T) { + svc := newTestBillingService() + + tests := []struct { + name string + models []string + input float64 + cacheRead float64 + output float64 + }{ + { + name: "Grok 4.3 family", + models: []string{ + "grok-4.3", + "grok-4.20-0309-reasoning", + "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent-0309", + "grok-4.20-reasoning", + "grok-4.20-non-reasoning", + }, + input: 1.25e-6, + cacheRead: 0.2e-6, + output: 2.5e-6, + }, + { + name: "Grok coding and Composer family", + models: []string{ + "grok-build", + "grok-build-0.1", + "grok-composer", + "grok-composer-2.5-fast", + "composer-2.5", + }, + input: 1e-6, + cacheRead: 0.2e-6, + output: 2e-6, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for _, model := range tt.models { + pricing, err := svc.GetModelPricing(model) + require.NoError(t, err, "model %s", model) + require.InDelta(t, tt.input, pricing.InputPricePerToken, 1e-12, "model %s input", model) + require.InDelta(t, tt.cacheRead, pricing.CacheReadPricePerToken, 1e-12, "model %s cached input", model) + require.InDelta(t, tt.output, pricing.OutputPricePerToken, 1e-12, "model %s output", model) + } + }) + } +} + func TestCalculateCost_SupportsCacheBreakdown(t *testing.T) { svc := &BillingService{ cfg: &config.Config{}, diff --git a/backend/internal/service/channel_monitor_checker.go b/backend/internal/service/channel_monitor_checker.go index 7fb829a3cb..ad4058f9e6 100644 --- a/backend/internal/service/channel_monitor_checker.go +++ b/backend/internal/service/channel_monitor_checker.go @@ -13,6 +13,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/tidwall/gjson" ) @@ -34,7 +35,7 @@ func newSSRFSafeHTTPClient(timeout time.Duration) *http.Client { TLSHandshakeTimeout: monitorTLSHandshakeTimeout, ResponseHeaderTimeout: monitorResponseHeaderTimeout, } - return &http.Client{Timeout: timeout, Transport: tr} + return &http.Client{Timeout: timeout, Transport: servertiming.WrapRoundTripper(tr)} } // CheckOptions 承载一次检测的自定义入参。 @@ -167,6 +168,7 @@ type providerAdapter struct { //nolint:gochecknoglobals // 适配器表是只读静态数据,初始化后不变更。 var providerAdapters = map[string]providerAdapter{ MonitorProviderOpenAI: providerOpenAIChatAdapter, + MonitorProviderGrok: providerGrokChatAdapter, MonitorProviderAnthropic: { buildPath: func(string) string { return providerAnthropicPath }, buildBody: func(model, prompt string) ([]byte, error) { @@ -204,20 +206,27 @@ var providerAdapters = map[string]providerAdapter{ } //nolint:gochecknoglobals // 适配器表是只读静态数据,初始化后不变更。 -var providerOpenAIChatAdapter = providerAdapter{ - buildPath: func(string) string { return providerOpenAIPath }, - buildBody: func(model, prompt string) ([]byte, error) { - return json.Marshal(map[string]any{ - "model": model, - "messages": []map[string]string{{"role": "user", "content": prompt}}, - "max_tokens": monitorChallengeMaxTokens, - "stream": false, - }) - }, - buildHeaders: func(apiKey string) map[string]string { - return map[string]string{"Authorization": "Bearer " + apiKey} - }, - textPath: "choices.0.message.content", +var providerOpenAIChatAdapter = newOpenAICompatibleChatAdapter(providerOpenAIPath) + +//nolint:gochecknoglobals // 适配器表是只读静态数据,初始化后不变更。 +var providerGrokChatAdapter = newOpenAICompatibleChatAdapter(providerGrokPath) + +func newOpenAICompatibleChatAdapter(path string) providerAdapter { + return providerAdapter{ + buildPath: func(string) string { return path }, + buildBody: func(model, prompt string) ([]byte, error) { + return json.Marshal(map[string]any{ + "model": model, + "messages": []map[string]string{{"role": "user", "content": prompt}}, + "max_tokens": monitorChallengeMaxTokens, + "stream": false, + }) + }, + buildHeaders: func(apiKey string) map[string]string { + return map[string]string{"Authorization": "Bearer " + apiKey} + }, + textPath: "choices.0.message.content", + } } //nolint:gochecknoglobals // 适配器表是只读静态数据,初始化后不变更。 @@ -407,8 +416,9 @@ func buildRequestBody(adapter providerAdapter, provider, apiMode, model, prompt var bodyMergeKeyDenyList = map[string]map[string]bool{ MonitorProviderOpenAI + ":" + MonitorAPIModeChatCompletions: {"model": true, "messages": true, "stream": true}, MonitorProviderOpenAI + ":" + MonitorAPIModeResponses: {"model": true, "instructions": true, "input": true, "stream": true}, - MonitorProviderAnthropic: {"model": true, "messages": true}, - MonitorProviderGemini: {"contents": true}, + MonitorProviderGrok: {"model": true, "messages": true, "stream": true}, + MonitorProviderAnthropic: {"model": true, "messages": true}, + MonitorProviderGemini: {"contents": true}, } func checkAPIMode(opts *CheckOptions) string { @@ -426,7 +436,7 @@ func bodyMergeDenyKey(provider, apiMode string) string { } func validateReplaceRequestBody(provider, apiMode string, body map[string]any) error { - if provider != MonitorProviderOpenAI { + if provider != MonitorProviderOpenAI && provider != MonitorProviderGrok { return nil } switch defaultAPIMode(apiMode) { @@ -527,6 +537,8 @@ var monitorAPIKeyPatterns = []struct { {regexp.MustCompile(`sk-ant-[A-Za-z0-9_-]{20,}`), "sk-ant-***REDACTED***"}, // OpenAI / Anthropic 通用 sk-: sk-xxxxxxx {regexp.MustCompile(`sk-[A-Za-z0-9-]{20,}`), "sk-***REDACTED***"}, + // xAI API Key:xai-xxxxxxx + {regexp.MustCompile(`xai-[A-Za-z0-9_-]{6,}`), "xai-***REDACTED***"}, // Gemini / Google API Key:固定前缀 + 35 位 {regexp.MustCompile(`AIza[A-Za-z0-9_-]{35}`), "AIza***REDACTED***"}, // JWT 三段式(Bearer 后常出现):eyJxxx.eyJxxx.signature @@ -536,7 +548,7 @@ var monitorAPIKeyPatterns = []struct { // sanitizeErrorMessage 擦除错误/响应文本中可能泄露的 API key。 // 处理两类来源: // 1. URL query 中的 ?key= / ?api_key= 等(Go *url.Error 会回填完整 URL) -// 2. 上游 HTTP body 文本里直接出现的 sk-* / AIza* / JWT 等密钥碎片 +// 2. 上游 HTTP body 文本里直接出现的 sk-* / xai-* / AIza* / JWT 等密钥碎片 // // 注意:与 gemini_messages_compat_service.go 的 sanitizeUpstreamErrorMessage 关注点类似但参数集更广, // 监控模块独立维护,避免互相耦合。 diff --git a/backend/internal/service/channel_monitor_checker_body_test.go b/backend/internal/service/channel_monitor_checker_body_test.go index bba3d7dfb7..bcf7af0b98 100644 --- a/backend/internal/service/channel_monitor_checker_body_test.go +++ b/backend/internal/service/channel_monitor_checker_body_test.go @@ -64,6 +64,7 @@ type openAICaptureHandler struct { lastHeaders http.Header lastPath string status int + rawResponse string responsesLeadingReasoning bool } @@ -80,6 +81,10 @@ func (h *openAICaptureHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) } w.Header().Set("Content-Type", "application/json") w.WriteHeader(h.status) + if h.rawResponse != "" { + _, _ = w.Write([]byte(h.rawResponse)) + return + } answer := answerFromOpenAIRequest(parsed) if h.lastPath == providerOpenAIResponsesPath { @@ -190,6 +195,90 @@ func TestRunCheckForModel_OpenAI_DefaultChatRequest(t *testing.T) { } } +func TestGrokMonitorConfiguration(t *testing.T) { + if err := validateProvider(MonitorProviderGrok); err != nil { + t.Fatalf("grok provider should be supported: %v", err) + } + if got := normalizeMonitorPrimaryModel(MonitorProviderGrok, ""); got != MonitorDefaultGrokModel { + t.Fatalf("expected default Grok model %q, got %q", MonitorDefaultGrokModel, got) + } + if err := validateAPIMode(MonitorProviderGrok, MonitorAPIModeChatCompletions); err != nil { + t.Fatalf("grok chat_completions mode should be valid: %v", err) + } + if err := validateAPIMode(MonitorProviderGrok, MonitorAPIModeResponses); err == nil { + t.Fatal("grok responses mode should be rejected by channel monitoring") + } + if err := validateReplaceRequestBody(MonitorProviderGrok, MonitorAPIModeChatCompletions, map[string]any{}); err == nil { + t.Fatal("grok replace-mode body should require messages") + } +} + +func TestRunCheckForModel_Grok_DefaultChatRequest(t *testing.T) { + h := &openAICaptureHandler{} + endpoint := setupFakeOpenAI(t, h) + + res := runCheckForModel(context.Background(), MonitorProviderGrok, endpoint, "xai-key", MonitorDefaultGrokModel, nil) + + if res.Status != MonitorStatusOperational { + t.Fatalf("Grok request should pass challenge, got status=%s message=%q", res.Status, res.Message) + } + if res.LatencyMs == nil { + t.Fatal("Grok request should record latency") + } + if h.lastPath != providerGrokPath { + t.Fatalf("expected Grok chat completions path %q, got %q", providerGrokPath, h.lastPath) + } + if h.lastBody["model"] != MonitorDefaultGrokModel { + t.Errorf("Grok body should contain model=%s, got %v", MonitorDefaultGrokModel, h.lastBody["model"]) + } + if _, ok := h.lastBody["messages"]; !ok { + t.Error("Grok body should contain messages") + } + if h.lastBody["stream"] != false { + t.Errorf("Grok body should set stream=false, got %v", h.lastBody["stream"]) + } + if h.lastHeaders.Get("Authorization") != "Bearer xai-key" { + t.Errorf("expected Grok bearer auth header, got %q", h.lastHeaders.Get("Authorization")) + } +} + +func TestRunCheckForModel_Grok_UpstreamFailure(t *testing.T) { + h := &openAICaptureHandler{status: http.StatusTooManyRequests} + endpoint := setupFakeOpenAI(t, h) + + res := runCheckForModel(context.Background(), MonitorProviderGrok, endpoint, "xai-key", MonitorDefaultGrokModel, nil) + + if res.Status != MonitorStatusError { + t.Fatalf("Grok 429 should be recorded as error, got status=%s message=%q", res.Status, res.Message) + } + if !strings.Contains(res.Message, "upstream HTTP 429") { + t.Fatalf("Grok failure should preserve upstream status, got %q", res.Message) + } + if res.LatencyMs == nil { + t.Fatal("Grok failure should still record latency") + } +} + +func TestRunCheckForModel_Grok_RedactsXAIKeyFromUpstreamBody(t *testing.T) { + h := &openAICaptureHandler{ + status: http.StatusUnauthorized, + rawResponse: `{"error":{"message":"invalid API key xai-secret"}}`, + } + endpoint := setupFakeOpenAI(t, h) + + res := runCheckForModel(context.Background(), MonitorProviderGrok, endpoint, "request-key", MonitorDefaultGrokModel, nil) + + if res.Status != MonitorStatusError { + t.Fatalf("Grok upstream failure should be recorded as error, got %s", res.Status) + } + if strings.Contains(res.Message, "xai-secret") { + t.Fatalf("Grok error message leaked xAI key: %q", res.Message) + } + if !strings.Contains(res.Message, "xai-***REDACTED***") { + t.Fatalf("Grok error message should contain redaction marker, got %q", res.Message) + } +} + func TestRunCheckForModel_OpenAIResponses_DefaultRequest(t *testing.T) { h := &openAICaptureHandler{} endpoint := setupFakeOpenAI(t, h) diff --git a/backend/internal/service/channel_monitor_const.go b/backend/internal/service/channel_monitor_const.go index 61f9f79894..2ee8eabde8 100644 --- a/backend/internal/service/channel_monitor_const.go +++ b/backend/internal/service/channel_monitor_const.go @@ -47,6 +47,8 @@ const ( // providerOpenAIPath OpenAI Chat Completions 路径。 providerOpenAIPath = "/v1/chat/completions" + // providerGrokPath Grok OpenAI-compatible Chat Completions 路径。 + providerGrokPath = "/v1/chat/completions" // providerOpenAIResponsesPath OpenAI Responses API 路径。 providerOpenAIResponsesPath = "/v1/responses" // providerAnthropicPath Anthropic Messages 路径。 @@ -54,10 +56,14 @@ const ( // providerGeminiPathTemplate Gemini generateContent 路径模板(含 model 占位)。 providerGeminiPathTemplate = "/v1beta/models/%s:generateContent" - // MonitorProviderOpenAI / Anthropic / Gemini provider 字符串常量(也是 ent enum 的实际值)。 + // MonitorProviderOpenAI / Anthropic / Gemini / Grok provider 字符串常量(也是 ent enum 的实际值)。 MonitorProviderOpenAI = "openai" MonitorProviderAnthropic = "anthropic" MonitorProviderGemini = "gemini" + MonitorProviderGrok = "grok" + + // MonitorDefaultGrokModel 是新增 Grok 监控未显式指定模型时使用的轻量测活模型。 + MonitorDefaultGrokModel = "grok-4.5" // MonitorStatusOperational 等监控状态字符串常量(与 ent enum 一致)。 MonitorStatusOperational = "operational" @@ -112,13 +118,13 @@ var ( "CHANNEL_MONITOR_NOT_FOUND", "channel monitor not found", ) ErrChannelMonitorInvalidProvider = infraerrors.BadRequest( - "CHANNEL_MONITOR_INVALID_PROVIDER", "provider must be one of openai/anthropic/gemini", + "CHANNEL_MONITOR_INVALID_PROVIDER", "provider must be one of openai/anthropic/gemini/grok", ) ErrChannelMonitorInvalidAPIMode = infraerrors.BadRequest( "CHANNEL_MONITOR_INVALID_API_MODE", "api_mode must be chat_completions or responses; responses is only supported for openai", ) ErrChannelMonitorInvalidRequestBody = infraerrors.BadRequest( - "CHANNEL_MONITOR_INVALID_REQUEST_BODY", "openai replace-mode body_override must include non-empty messages for chat_completions or non-empty instructions and input for responses", + "CHANNEL_MONITOR_INVALID_REQUEST_BODY", "openai-compatible replace-mode body_override must include non-empty messages for chat_completions or non-empty instructions and input for responses", ) ErrChannelMonitorInvalidInterval = infraerrors.BadRequest( "CHANNEL_MONITOR_INVALID_INTERVAL", "interval_seconds must be in [15, 3600]", diff --git a/backend/internal/service/channel_monitor_service.go b/backend/internal/service/channel_monitor_service.go index 7b53bb20b0..b5dea22589 100644 --- a/backend/internal/service/channel_monitor_service.go +++ b/backend/internal/service/channel_monitor_service.go @@ -123,7 +123,7 @@ func (s *ChannelMonitorService) Create(ctx context.Context, p ChannelMonitorCrea APIMode: defaultAPIMode(p.APIMode), Endpoint: normalizeEndpoint(p.Endpoint), APIKey: encrypted, // 注意:传入 repository 时该字段为密文 - PrimaryModel: strings.TrimSpace(p.PrimaryModel), + PrimaryModel: normalizeMonitorPrimaryModel(p.Provider, p.PrimaryModel), ExtraModels: normalizeModels(p.ExtraModels), GroupName: strings.TrimSpace(p.GroupName), Enabled: p.Enabled, @@ -167,7 +167,7 @@ func validateCreateParams(p ChannelMonitorCreateParams) error { if strings.TrimSpace(p.APIKey) == "" { return ErrChannelMonitorMissingAPIKey } - if strings.TrimSpace(p.PrimaryModel) == "" { + if normalizeMonitorPrimaryModel(p.Provider, p.PrimaryModel) == "" { return ErrChannelMonitorMissingPrimaryModel } return nil @@ -486,8 +486,8 @@ func applyMonitorUpdate(existing *ChannelMonitor, p ChannelMonitorUpdateParams) if err := validateProvider(*p.Provider); err != nil { return err } + providerChanged = existing.Provider != *p.Provider existing.Provider = *p.Provider - providerChanged = true } if p.Endpoint != nil { if err := validateEndpoint(*p.Endpoint); err != nil { @@ -496,7 +496,13 @@ func applyMonitorUpdate(existing *ChannelMonitor, p ChannelMonitorUpdateParams) existing.Endpoint = normalizeEndpoint(*p.Endpoint) } if p.PrimaryModel != nil { - existing.PrimaryModel = strings.TrimSpace(*p.PrimaryModel) + primaryModel := normalizeMonitorPrimaryModel(existing.Provider, *p.PrimaryModel) + if primaryModel == "" { + return ErrChannelMonitorMissingPrimaryModel + } + existing.PrimaryModel = primaryModel + } else if providerChanged && existing.Provider == MonitorProviderGrok { + existing.PrimaryModel = MonitorDefaultGrokModel } if p.ExtraModels != nil { existing.ExtraModels = normalizeModels(*p.ExtraModels) diff --git a/backend/internal/service/channel_monitor_service_grok_test.go b/backend/internal/service/channel_monitor_service_grok_test.go new file mode 100644 index 0000000000..20c9db2666 --- /dev/null +++ b/backend/internal/service/channel_monitor_service_grok_test.go @@ -0,0 +1,85 @@ +//go:build unit + +package service + +import "testing" + +func TestApplyMonitorUpdate_ProviderOnlySwitchToGrokUsesDefaultModel(t *testing.T) { + grok := MonitorProviderGrok + existing := &ChannelMonitor{ + Provider: MonitorProviderOpenAI, + APIMode: MonitorAPIModeResponses, + PrimaryModel: "gpt-5", + IntervalSeconds: 60, + } + + err := applyMonitorUpdate(existing, ChannelMonitorUpdateParams{Provider: &grok}) + if err != nil { + t.Fatalf("provider-only switch to Grok failed: %v", err) + } + if existing.PrimaryModel != MonitorDefaultGrokModel { + t.Fatalf("expected Grok default model %q, got %q", MonitorDefaultGrokModel, existing.PrimaryModel) + } + if existing.APIMode != MonitorAPIModeChatCompletions { + t.Fatalf("expected Grok API mode %q, got %q", MonitorAPIModeChatCompletions, existing.APIMode) + } +} + +func TestApplyMonitorUpdate_SwitchToGrokPreservesExplicitModel(t *testing.T) { + grok := MonitorProviderGrok + explicitModel := "grok-4.3" + existing := &ChannelMonitor{ + Provider: MonitorProviderOpenAI, + APIMode: MonitorAPIModeChatCompletions, + PrimaryModel: "gpt-5", + IntervalSeconds: 60, + } + + err := applyMonitorUpdate(existing, ChannelMonitorUpdateParams{ + Provider: &grok, + PrimaryModel: &explicitModel, + }) + if err != nil { + t.Fatalf("switch to Grok with explicit model failed: %v", err) + } + if existing.PrimaryModel != explicitModel { + t.Fatalf("expected explicit model %q, got %q", explicitModel, existing.PrimaryModel) + } +} + +func TestApplyMonitorUpdate_SameGrokProviderDoesNotResetExistingModel(t *testing.T) { + grok := MonitorProviderGrok + existing := &ChannelMonitor{ + Provider: MonitorProviderGrok, + APIMode: MonitorAPIModeChatCompletions, + PrimaryModel: "grok-4.3", + IntervalSeconds: 60, + } + + err := applyMonitorUpdate(existing, ChannelMonitorUpdateParams{Provider: &grok}) + if err != nil { + t.Fatalf("same-provider Grok update failed: %v", err) + } + if existing.PrimaryModel != "grok-4.3" { + t.Fatalf("same-provider update reset existing model to %q", existing.PrimaryModel) + } +} + +func TestApplyMonitorUpdate_SwitchToGrokRejectsResponsesMode(t *testing.T) { + grok := MonitorProviderGrok + responses := MonitorAPIModeResponses + existing := &ChannelMonitor{ + Provider: MonitorProviderOpenAI, + APIMode: MonitorAPIModeChatCompletions, + PrimaryModel: "gpt-5", + IntervalSeconds: 60, + } + + err := applyMonitorUpdate(existing, ChannelMonitorUpdateParams{ + Provider: &grok, + APIMode: &responses, + }) + if err == nil { + t.Fatal("Grok responses mode should remain unsupported") + } +} diff --git a/backend/internal/service/channel_monitor_template_types.go b/backend/internal/service/channel_monitor_template_types.go index 03cd518d28..0b824d577d 100644 --- a/backend/internal/service/channel_monitor_template_types.go +++ b/backend/internal/service/channel_monitor_template_types.go @@ -55,7 +55,7 @@ var ( "CHANNEL_MONITOR_TEMPLATE_NOT_FOUND", "channel monitor request template not found", ) ErrChannelMonitorTemplateInvalidProvider = infraerrors.BadRequest( - "CHANNEL_MONITOR_TEMPLATE_INVALID_PROVIDER", "template provider must be one of openai/anthropic/gemini", + "CHANNEL_MONITOR_TEMPLATE_INVALID_PROVIDER", "template provider must be one of openai/anthropic/gemini/grok", ) ErrChannelMonitorTemplateInvalidAPIMode = infraerrors.BadRequest( "CHANNEL_MONITOR_TEMPLATE_INVALID_API_MODE", "template api_mode must be chat_completions or responses; responses is only supported for openai", diff --git a/backend/internal/service/channel_monitor_validate.go b/backend/internal/service/channel_monitor_validate.go index c5a4783b91..7740dc83b7 100644 --- a/backend/internal/service/channel_monitor_validate.go +++ b/backend/internal/service/channel_monitor_validate.go @@ -124,6 +124,16 @@ func normalizeModels(in []string) []string { return out } +// normalizeMonitorPrimaryModel applies the Grok health-check default while +// preserving the existing required-model behavior for every other provider. +func normalizeMonitorPrimaryModel(provider, model string) string { + model = strings.TrimSpace(model) + if model == "" && provider == MonitorProviderGrok { + return MonitorDefaultGrokModel + } + return model +} + // defaultAPIMode 空串归一为 chat_completions,保证历史数据与旧客户端兼容。 func defaultAPIMode(apiMode string) string { if strings.TrimSpace(apiMode) == "" { diff --git a/backend/internal/service/concurrency_service.go b/backend/internal/service/concurrency_service.go index f2f2aade89..df18379704 100644 --- a/backend/internal/service/concurrency_service.go +++ b/backend/internal/service/concurrency_service.go @@ -6,6 +6,7 @@ import ( "crypto/sha256" "encoding/binary" "encoding/hex" + "errors" "os" "strconv" "sync" @@ -13,6 +14,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" "golang.org/x/sync/singleflight" ) @@ -59,6 +61,131 @@ type APIKeyConcurrencyCache interface { GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) } +// OpenAIWSIngressLeaseCache owns the short-lived distributed lease used to +// bound live client WebSocket sessions. It is deliberately independent of the +// request-slot namespace: idle ingress connections do not occupy turn slots. +type OpenAIWSIngressLeaseCache interface { + AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) + RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error) + ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error +} + +const ( + openAIWSIngressLeaseTTL = 60 * time.Second + openAIWSIngressLeaseRefreshInterval = 20 * time.Second + openAIWSIngressLeaseOperationTO = 2 * time.Second +) + +var ErrOpenAIWSIngressLeaseLost = errors.New("openai websocket ingress lease lost") + +// OpenAIWSIngressLease keeps a Redis-backed ingress lease alive and cancels +// its context if Redis cannot confirm ownership for a full lease lifetime. +// Call Release on every handler exit to reclaim capacity immediately. +type OpenAIWSIngressLease struct { + ctx context.Context + cancel context.CancelCauseFunc + cache OpenAIWSIngressLeaseCache + apiKeyID int64 + leaseID string + + stopOnce sync.Once + stopCh chan struct{} + refreshDone chan struct{} +} + +func (l *OpenAIWSIngressLease) Context() context.Context { + if l == nil || l.ctx == nil { + return context.Background() + } + return l.ctx +} + +func (l *OpenAIWSIngressLease) Release() { + if l == nil { + return + } + l.stopOnce.Do(func() { + if l.stopCh != nil { + close(l.stopCh) + } + if l.cancel != nil { + l.cancel(nil) + } + if l.refreshDone != nil { + <-l.refreshDone + } + if l.cache == nil || l.apiKeyID <= 0 || l.leaseID == "" { + return + } + releaseCtx, releaseCancel := context.WithTimeout(context.Background(), openAIWSIngressLeaseOperationTO) + defer releaseCancel() + if err := l.cache.ReleaseOpenAIWSIngressLease(releaseCtx, l.apiKeyID, l.leaseID); err != nil { + logger.L().Warn("openai_ws_ingress_lease_release_failed", + zap.Int64("api_key_id", l.apiKeyID), + zap.Error(err), + ) + } + }) +} + +func (l *OpenAIWSIngressLease) refreshLoop() { + defer func() { + if l != nil && l.refreshDone != nil { + close(l.refreshDone) + } + }() + if l == nil || l.cache == nil { + return + } + ticker := time.NewTicker(openAIWSIngressLeaseRefreshInterval) + defer ticker.Stop() + lastConfirmedAt := time.Now() + for { + select { + case <-l.ctx.Done(): + return + case <-l.stopCh: + return + case <-ticker.C: + var lost bool + lastConfirmedAt, lost = l.refresh(lastConfirmedAt) + if lost { + l.cancel(ErrOpenAIWSIngressLeaseLost) + return + } + } + } +} + +// refresh confirms the lease is still owned. A missing member is an immediate +// lease loss; transient Redis errors are tolerated only for one full lease TTL. +func (l *OpenAIWSIngressLease) refresh(lastConfirmedAt time.Time) (time.Time, bool) { + refreshCtx, refreshCancel := context.WithTimeout(context.Background(), openAIWSIngressLeaseOperationTO) + owned, err := l.cache.RefreshOpenAIWSIngressLease(refreshCtx, l.apiKeyID, l.leaseID) + refreshCancel() + if err == nil && owned { + return time.Now(), false + } + if err == nil { + err = ErrOpenAIWSIngressLeaseLost + } + elapsed := time.Since(lastConfirmedAt) + logger.L().Warn("openai_ws_ingress_lease_refresh_failed", + zap.Int64("api_key_id", l.apiKeyID), + zap.Duration("unconfirmed_for", elapsed), + zap.Error(err), + ) + if errors.Is(err, ErrOpenAIWSIngressLeaseLost) || elapsed >= openAIWSIngressLeaseTTL { + logger.L().Error("openai_ws_ingress_lease_lost", + zap.Int64("api_key_id", l.apiKeyID), + zap.Duration("unconfirmed_for", elapsed), + zap.Error(err), + ) + return lastConfirmedAt, true + } + return lastConfirmedAt, false +} + var ( requestIDPrefix = initRequestIDPrefix() requestIDCounter atomic.Uint64 @@ -125,6 +252,47 @@ func NewConcurrencyService(cache ConcurrencyCache) *ConcurrencyService { return svc } +// AcquireOpenAIWSIngressLease atomically reserves one live ingress connection +// for an API key. A non-positive limit explicitly disables this protection. +func (s *ConcurrencyService) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int) (*OpenAIWSIngressLease, bool, error) { + if maxConnections <= 0 { + return nil, true, nil + } + if s == nil || s.cache == nil || apiKeyID <= 0 { + return nil, false, errors.New("openai websocket ingress lease cache is unavailable") + } + cache, ok := s.cache.(OpenAIWSIngressLeaseCache) + if !ok { + return nil, false, errors.New("openai websocket ingress lease cache is unsupported") + } + leaseID := generateRequestID() + baseCtx := context.Background() + if ctx != nil { + baseCtx = context.WithoutCancel(ctx) + } + acquireCtx, acquireCancel := context.WithTimeout(baseCtx, openAIWSIngressLeaseOperationTO) + acquired, err := cache.AcquireOpenAIWSIngressLease(acquireCtx, apiKeyID, maxConnections, leaseID) + acquireCancel() + if err != nil || !acquired { + return nil, acquired, err + } + if ctx == nil { + ctx = context.Background() + } + leaseCtx, leaseCancel := context.WithCancelCause(ctx) + lease := &OpenAIWSIngressLease{ + ctx: leaseCtx, + cancel: leaseCancel, + cache: cache, + apiKeyID: apiKeyID, + leaseID: leaseID, + stopCh: make(chan struct{}), + refreshDone: make(chan struct{}), + } + go lease.refreshLoop() + return lease, true, nil +} + // SetAccountLoadBatchCacheTTL 设置账号负载批量读取的极短 TTL 缓存;非正数表示禁用缓存。 func (s *ConcurrencyService) SetAccountLoadBatchCacheTTL(ttl time.Duration) { if s == nil { diff --git a/backend/internal/service/concurrency_service_test.go b/backend/internal/service/concurrency_service_test.go index 3f358bbe6a..d079c7e60a 100644 --- a/backend/internal/service/concurrency_service_test.go +++ b/backend/internal/service/concurrency_service_test.go @@ -45,7 +45,47 @@ type stubConcurrencyCacheForTest struct { releasedAPIKeyRequestIDs []string } +type ingressLeaseCacheForTest struct { + stubConcurrencyCacheForTest + acquireIngressResult bool + acquireIngressErr error + acquireIngressFn func(context.Context, int64, int, string) (bool, error) + refreshIngressResult bool + refreshIngressErr error + refreshIngressFn func(context.Context, int64, string) (bool, error) + releaseIngressErr error + releaseIngressFn func(context.Context, int64, string) error + acquireIngressCalls int + refreshIngressCalls int + releaseIngressCalls int +} + +func (c *ingressLeaseCacheForTest) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) { + c.acquireIngressCalls++ + if c.acquireIngressFn != nil { + return c.acquireIngressFn(ctx, apiKeyID, maxConnections, leaseID) + } + return c.acquireIngressResult, c.acquireIngressErr +} + +func (c *ingressLeaseCacheForTest) RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error) { + c.refreshIngressCalls++ + if c.refreshIngressFn != nil { + return c.refreshIngressFn(ctx, apiKeyID, leaseID) + } + return c.refreshIngressResult, c.refreshIngressErr +} + +func (c *ingressLeaseCacheForTest) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error { + c.releaseIngressCalls++ + if c.releaseIngressFn != nil { + return c.releaseIngressFn(ctx, apiKeyID, leaseID) + } + return c.releaseIngressErr +} + var _ ConcurrencyCache = (*stubConcurrencyCacheForTest)(nil) +var _ OpenAIWSIngressLeaseCache = (*ingressLeaseCacheForTest)(nil) func (c *stubConcurrencyCacheForTest) AcquireAccountSlot(_ context.Context, _ int64, _ int, _ string) (bool, error) { return c.acquireResult, c.acquireErr @@ -285,6 +325,114 @@ func TestGetAPIKeyConcurrencyBatch_Fallbacks(t *testing.T) { }) } +func TestAcquireOpenAIWSIngressLease(t *testing.T) { + t.Run("zero value release is safe", func(t *testing.T) { + var lease OpenAIWSIngressLease + require.NotPanics(t, lease.Release) + }) + + t.Run("disabled", func(t *testing.T) { + cache := &ingressLeaseCacheForTest{} + lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(nil, 1, 0) + require.NoError(t, err) + require.True(t, acquired) + require.Nil(t, lease) + require.Zero(t, cache.acquireIngressCalls) + }) + + t.Run("unsupported cache fails closed", func(t *testing.T) { + lease, acquired, err := NewConcurrencyService(&stubConcurrencyCacheForTest{}).AcquireOpenAIWSIngressLease(context.Background(), 1, 1) + require.Error(t, err) + require.False(t, acquired) + require.Nil(t, lease) + }) + + t.Run("capacity rejected", func(t *testing.T) { + cache := &ingressLeaseCacheForTest{acquireIngressResult: false} + lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(context.Background(), 1, 1) + require.NoError(t, err) + require.False(t, acquired) + require.Nil(t, lease) + }) + + t.Run("release returns capacity", func(t *testing.T) { + cache := &ingressLeaseCacheForTest{acquireIngressResult: true, refreshIngressResult: true} + lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(nil, 1, 1) + require.NoError(t, err) + require.True(t, acquired) + require.NotNil(t, lease) + lease.Release() + lease.Release() + require.Equal(t, 1, cache.releaseIngressCalls) + }) +} + +func TestOpenAIWSIngressLeaseRefreshLoss(t *testing.T) { + t.Run("missing lease is lost immediately", func(t *testing.T) { + cache := &ingressLeaseCacheForTest{refreshIngressResult: false} + lease := &OpenAIWSIngressLease{cache: cache, apiKeyID: 1, leaseID: "missing"} + _, lost := lease.refresh(time.Now()) + require.True(t, lost) + require.Equal(t, 1, cache.refreshIngressCalls) + }) + + t.Run("persistent redis errors lose lease after ttl", func(t *testing.T) { + cache := &ingressLeaseCacheForTest{refreshIngressErr: errors.New("redis unavailable")} + lease := &OpenAIWSIngressLease{cache: cache, apiKeyID: 1, leaseID: "unconfirmed"} + _, lost := lease.refresh(time.Now().Add(-openAIWSIngressLeaseTTL)) + require.True(t, lost) + require.Equal(t, 1, cache.refreshIngressCalls) + }) +} + +func TestOpenAIWSIngressLeaseReleaseWaitsForInFlightRefresh(t *testing.T) { + refreshStarted := make(chan struct{}) + allowRefresh := make(chan struct{}) + cache := &ingressLeaseCacheForTest{ + refreshIngressFn: func(context.Context, int64, string) (bool, error) { + close(refreshStarted) + <-allowRefresh + return true, nil + }, + } + ctx, cancel := context.WithCancelCause(context.Background()) + lease := &OpenAIWSIngressLease{ + ctx: ctx, + cancel: cancel, + cache: cache, + apiKeyID: 1, + leaseID: "in-flight-refresh", + stopCh: make(chan struct{}), + refreshDone: make(chan struct{}), + } + go func() { + defer close(lease.refreshDone) + _, _ = lease.refresh(time.Now()) + }() + <-refreshStarted + + released := make(chan struct{}) + go func() { + lease.Release() + close(released) + }() + + select { + case <-released: + t.Fatal("release returned before the in-flight refresh completed") + case <-time.After(20 * time.Millisecond): + } + require.Zero(t, cache.releaseIngressCalls) + + close(allowRefresh) + select { + case <-released: + case <-time.After(time.Second): + t.Fatal("release did not complete after the refresh returned") + } + require.Equal(t, 1, cache.releaseIngressCalls) +} + func TestGenerateRequestID_UsesStablePrefixAndMonotonicCounter(t *testing.T) { id1 := generateRequestID() id2 := generateRequestID() diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index 6d3b91d205..f633c8ad17 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -22,6 +22,7 @@ import ( infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) const ( @@ -561,7 +562,7 @@ func NewContentModerationService( userRepo: userRepo, authCacheInvalidator: authCacheInvalidator, emailService: emailService, - httpClient: &http.Client{}, + httpClient: servertiming.InstrumentClient(nil), workerCount: maxContentModerationWorkerCount, asyncQueue: make(chan contentModerationTask, maxContentModerationQueueSize), keyHealth: make(map[string]*contentModerationKeyHealth), diff --git a/backend/internal/service/crs_sync_long_context_billing_test.go b/backend/internal/service/crs_sync_long_context_billing_test.go new file mode 100644 index 0000000000..6439f08190 --- /dev/null +++ b/backend/internal/service/crs_sync_long_context_billing_test.go @@ -0,0 +1,169 @@ +//go:build unit + +package service + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type crsLongContextAccountRepo struct { + AccountRepository + accounts map[string]*Account + nextID int64 +} + +type crsOpenAILongContextSource struct { + collection string + credentials map[string]any + extra map[string]any +} + +func newCRSLongContextAccountRepo(existing ...*Account) *crsLongContextAccountRepo { + repo := &crsLongContextAccountRepo{accounts: make(map[string]*Account)} + for _, account := range existing { + if account == nil { + continue + } + crsID, _ := account.Extra["crs_account_id"].(string) + repo.accounts[crsID] = account + if account.ID > repo.nextID { + repo.nextID = account.ID + } + } + return repo +} + +func (r *crsLongContextAccountRepo) Create(_ context.Context, account *Account) error { + r.nextID++ + account.ID = r.nextID + crsID, _ := account.Extra["crs_account_id"].(string) + r.accounts[crsID] = account + return nil +} + +func (r *crsLongContextAccountRepo) Update(_ context.Context, account *Account) error { + crsID, _ := account.Extra["crs_account_id"].(string) + r.accounts[crsID] = account + return nil +} + +func (r *crsLongContextAccountRepo) GetByCRSAccountID(_ context.Context, crsID string) (*Account, error) { + return r.accounts[crsID], nil +} + +func (r *crsLongContextAccountRepo) ListShadowsByParent(_ context.Context, _ int64) ([]*Account, error) { + return nil, nil +} + +func TestCRSSyncOpenAILongContextBilling(t *testing.T) { + tests := []struct { + name string + collection string + credentials map[string]any + sourceExtra map[string]any + existingExtra map[string]any + wantAction string + wantEnabled bool + }{ + {name: "OAuth create defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, wantAction: "created"}, + {name: "OAuth create preserves source true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true}, + {name: "OAuth create preserves source false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"}, + {name: "OAuth update defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"}, + {name: "OAuth update preserves existing true when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true}, + {name: "OAuth update preserves existing false when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"}, + {name: "OAuth update preserves source true over existing false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true}, + {name: "OAuth update preserves source false over existing true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"}, + {name: "OAuth rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "OAuth rejects malformed existing value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "OAuth update rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"}, + {name: "API key create defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, wantAction: "created"}, + {name: "API key create preserves source true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true}, + {name: "API key create preserves source false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"}, + {name: "API key update defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"}, + {name: "API key update preserves existing true when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true}, + {name: "API key update preserves existing false when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"}, + {name: "API key update preserves source true over existing false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true}, + {name: "API key update preserves source false over existing true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"}, + {name: "API key rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "API key rejects malformed existing value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "API key update rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + const crsID = "crs-openai-1" + var existing *Account + if tt.existingExtra != nil { + existingExtra := mergeMap(tt.existingExtra, map[string]any{"crs_account_id": crsID}) + accountType := AccountTypeOAuth + if tt.collection == "openaiResponsesAccounts" { + accountType = AccountTypeAPIKey + } + existing = &Account{ID: 41, Platform: PlatformOpenAI, Type: accountType, Extra: existingExtra} + } + repo := newCRSLongContextAccountRepo(existing) + result := runCRSOpenAILongContextSync(t, repo, crsOpenAILongContextSource{ + collection: tt.collection, + credentials: tt.credentials, + extra: tt.sourceExtra, + }) + + require.Len(t, result.Items, 1) + require.Equal(t, tt.wantAction, result.Items[0].Action) + if tt.wantAction == "failed" { + require.Contains(t, result.Items[0].Error, "openai_long_context_billing_enabled must be a boolean") + return + } + stored, ok := repo.accounts[crsID].Extra[openAILongContextBillingEnabledKey] + require.True(t, ok) + require.Equal(t, tt.wantEnabled, stored) + }) + } +} + +func runCRSOpenAILongContextSync(t *testing.T, repo AccountRepository, source crsOpenAILongContextSource) *SyncFromCRSResult { + t.Helper() + account := map[string]any{ + "kind": "openai", + "id": "crs-openai-1", + "name": "OpenAI CRS", + "isActive": true, + "schedulable": true, + "credentials": source.credentials, + } + if source.extra != nil { + account["extra"] = source.extra + } + + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + response.Header().Set("Content-Type", "application/json") + if request.URL.Path == "/web/auth/login" { + _, _ = response.Write([]byte(`{"success":true,"token":"admin-token"}`)) + return + } + require.Equal(t, "/admin/sync/export-accounts", request.URL.Path) + require.NoError(t, json.NewEncoder(response).Encode(map[string]any{ + "success": true, + "data": map[string]any{source.collection: []any{account}}, + })) + })) + t.Cleanup(server.Close) + + cfg := &config.Config{} + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + service := NewCRSSyncService(repo, nil, nil, nil, nil, cfg) + result, err := service.SyncFromCRS(context.Background(), SyncFromCRSInput{ + BaseURL: server.URL, + Username: "admin", + Password: "password", + }) + require.NoError(t, err) + return result +} diff --git a/backend/internal/service/crs_sync_service.go b/backend/internal/service/crs_sync_service.go index edf3cd43d2..d0abc74038 100644 --- a/backend/internal/service/crs_sync_service.go +++ b/backend/internal/service/crs_sync_service.go @@ -168,6 +168,7 @@ type crsOpenAIResponsesAccount struct { Status string `json:"status"` Proxy *crsProxy `json:"proxy"` Credentials map[string]any `json:"credentials"` + Extra map[string]any `json:"extra"` } type crsOpenAIOAuthAccount struct { @@ -632,6 +633,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + var existingExtra map[string]any + if existing != nil { + existingExtra = existing.Extra + } + extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra) + if err != nil { + item.Action = "failed" + item.Error = err.Error() + result.Failed++ + result.Items = append(result.Items, item) + continue + } if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -670,7 +683,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeOAuth @@ -751,11 +764,13 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput concurrency := 3 status := mapCRSStatus(src.IsActive, src.Status) - extra := map[string]any{ - "crs_account_id": src.ID, - "crs_kind": src.Kind, - "crs_synced_at": now, + extra := make(map[string]any, len(src.Extra)+3) + for key, value := range src.Extra { + extra[key] = value } + extra["crs_account_id"] = src.ID + extra["crs_kind"] = src.Kind + extra["crs_synced_at"] = now existing, err := s.accountRepo.GetByCRSAccountID(ctx, src.ID) if err != nil { @@ -765,6 +780,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + var existingExtra map[string]any + if existing != nil { + existingExtra = existing.Extra + } + extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra) + if err != nil { + item.Action = "failed" + item.Error = err.Error() + result.Failed++ + result.Items = append(result.Items, item) + continue + } if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -809,7 +836,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeAPIKey @@ -1098,6 +1125,10 @@ func mergeMap(existing map[string]any, updates map[string]any) map[string]any { return out } +func mergeCRSOpenAILongContextBillingExtra(existing, updates map[string]any) (map[string]any, error) { + return normalizeOpenAILongContextBillingExtra(PlatformOpenAI, mergeMap(existing, updates)) +} + func (s *CRSSyncService) mapOrCreateProxy(ctx context.Context, enabled bool, cached *[]Proxy, src *crsProxy, defaultName string) (*int64, error) { if !enabled || src == nil { return nil, nil diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index 8a95915981..61ab3abd2f 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -947,6 +947,7 @@ func (s *GatewayService) buildRecordUsageLog( usageLog.CacheReadCost = cost.CacheReadCost usageLog.TotalCost = cost.TotalCost usageLog.ActualCost = cost.ActualCost + usageLog.LongContextBillingApplied = cost.LongContextBillingApplied } return usageLog diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 154e3003ef..100d720659 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -26,6 +26,8 @@ const ( GrokMediaEndpointImagesGenerations GrokMediaEndpoint = "images_generations" GrokMediaEndpointImagesEdits GrokMediaEndpoint = "images_edits" GrokMediaEndpointVideosGenerations GrokMediaEndpoint = "videos_generations" + GrokMediaEndpointVideosEdits GrokMediaEndpoint = "videos_edits" + GrokMediaEndpointVideosExtensions GrokMediaEndpoint = "videos_extensions" GrokMediaEndpointVideoStatus GrokMediaEndpoint = "video_status" ) @@ -35,7 +37,7 @@ func (e GrokMediaEndpoint) RequiresRequestBody() bool { func (e GrokMediaEndpoint) IsGenerationRequest() bool { switch e { - case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits, GrokMediaEndpointVideosGenerations: + case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits, GrokMediaEndpointVideosGenerations, GrokMediaEndpointVideosEdits, GrokMediaEndpointVideosExtensions: return true default: return false @@ -274,6 +276,10 @@ func (e GrokMediaEndpoint) upstreamURL(baseURL, requestID string) (string, error return xai.BuildImagesEditsURL(baseURL) case GrokMediaEndpointVideosGenerations: return xai.BuildVideosGenerationsURL(baseURL) + case GrokMediaEndpointVideosEdits: + return xai.BuildVideosEditsURL(baseURL) + case GrokMediaEndpointVideosExtensions: + return xai.BuildVideosExtensionsURL(baseURL) case GrokMediaEndpointVideoStatus: return xai.BuildVideoURL(baseURL, requestID) default: @@ -302,7 +308,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( if err != nil { return nil, err } - targetURL, err := endpoint.upstreamURL(account.GetGrokBaseURL(), requestID) + targetURL, err := endpoint.upstreamURL(account.GetGrokMediaBaseURL(), requestID) if err != nil { return nil, err } @@ -333,7 +339,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( } upstreamReq.Header.Set("Authorization", "Bearer "+token) upstreamReq.Header.Set("Accept", "application/json") - upstreamReq.Header.Set("User-Agent", "sub2api-grok/1.0") + applyGrokCLIHeaders(upstreamReq.Header) if endpoint.RequiresRequestBody() { contentType = strings.TrimSpace(contentType) if contentType == "" { @@ -357,11 +363,10 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( requestIDHeader := firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")) requestModel := requestInfo.Model if resp.StatusCode >= 400 { - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) return s.handleGrokMediaErrorResponse(ctx, resp, c, account, requestIDHeader, requestModel) } - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) if err != nil { return nil, err @@ -532,7 +537,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi meta.ImageSize = requestInfo.SizeTier meta.ImageInputSize = requestInfo.Size meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody) - case GrokMediaEndpointVideosGenerations: + case GrokMediaEndpointVideosGenerations, GrokMediaEndpointVideosEdits, GrokMediaEndpointVideosExtensions: meta.ResponseID = extractGrokMediaVideoRequestID(responseBody) meta.VideoCount = 1 meta.VideoResolution = requestInfo.Resolution @@ -564,6 +569,9 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse( requestedModel string, ) (*OpenAIForwardResult, error) { body := s.readUpstreamErrorBody(resp) + // Reconcile readiness before configurable passthrough branches can return; + // otherwise a Grok 429 can remain schedulable. + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body) upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(body))) if upstreamMsg == "" { upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode) @@ -609,7 +617,6 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse( return nil, fmt.Errorf("upstream error: %d (not in custom error codes) message=%s", resp.StatusCode, upstreamMsg) } - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body) kind := "http_error" if s.shouldFailoverUpstreamError(resp.StatusCode) { kind = "failover" diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index acad39985f..136e5e3b21 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -3,8 +3,6 @@ package service import ( "context" "crypto/subtle" - "encoding/base64" - "encoding/json" "net/http" "strings" "time" @@ -101,6 +99,8 @@ type GrokTokenInfo struct { ClientID string `json:"client_id,omitempty"` Scope string `json:"scope,omitempty"` Email string `json:"email,omitempty"` + Subject string `json:"sub,omitempty"` + TeamID string `json:"team_id,omitempty"` SubscriptionTier string `json:"subscription_tier,omitempty"` EntitlementStatus string `json:"entitlement_status,omitempty"` } @@ -175,6 +175,18 @@ func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToke return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID()) } +func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { + proxyURL, err := s.proxyURL(ctx, proxyID) + if err != nil { + return nil, err + } + tokenResp, err := s.oauthClient.ConvertSSOToBuild(ctx, ssoToken, proxyURL) + if err != nil { + return nil, err + } + return s.tokenInfoFromResponse(tokenResp, xai.DefaultClientID, nil), nil +} + func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) { if account == nil || account.Platform != PlatformGrok { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT", "account is not a Grok account") @@ -229,13 +241,19 @@ func (s *GrokOAuthService) BuildAccountCredentials(tokenInfo *GrokTokenInfo) map if tokenInfo.Email != "" { creds["email"] = tokenInfo.Email } + if tokenInfo.Subject != "" { + creds["sub"] = tokenInfo.Subject + } + if tokenInfo.TeamID != "" { + creds["team_id"] = tokenInfo.TeamID + } if tokenInfo.SubscriptionTier != "" { creds["subscription_tier"] = tokenInfo.SubscriptionTier } if tokenInfo.EntitlementStatus != "" { creds["entitlement_status"] = tokenInfo.EntitlementStatus } - creds["base_url"] = xai.DefaultBaseURL + creds["base_url"] = xai.DefaultCLIBaseURL return creds } @@ -265,12 +283,23 @@ func (s *GrokOAuthService) tokenInfoFromResponse(tokenResp *xai.TokenResponse, c if info.TokenType == "" { info.TokenType = "Bearer" } - if email := parseJWTEmailClaim(tokenResp.IDToken); email != "" { - info.Email = email - } - if info.Email == "" && existing != nil { - if email, _ := existing["email"].(string); email != "" { - info.Email = email + applyGrokTokenClaims(info, tokenResp.IDToken) + applyGrokTokenClaims(info, tokenResp.AccessToken) + if existing != nil { + if info.Email == "" { + if email, _ := existing["email"].(string); email != "" { + info.Email = email + } + } + if info.Subject == "" { + if subject, _ := existing["sub"].(string); subject != "" { + info.Subject = subject + } + } + if info.TeamID == "" { + if teamID, _ := existing["team_id"].(string); teamID != "" { + info.TeamID = teamID + } } } return info @@ -293,20 +322,21 @@ func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string return proxy.URL(), nil } -func parseJWTEmailClaim(token string) string { - parts := strings.Split(token, ".") - if len(parts) < 2 { - return "" +func applyGrokTokenClaims(info *GrokTokenInfo, token string) { + if info == nil || strings.TrimSpace(token) == "" { + return } - payload, err := base64.RawURLEncoding.DecodeString(parts[1]) - if err != nil { - return "" + claims := xai.DecodeJWTClaims(token) + if claims == nil { + return } - var claims struct { - Email string `json:"email"` + if info.Email == "" { + info.Email = xai.JWTClaimString(claims, "email") } - if err := json.Unmarshal(payload, &claims); err != nil { - return "" + if info.Subject == "" { + info.Subject = xai.JWTClaimString(claims, "sub") + } + if info.TeamID == "" { + info.TeamID = xai.JWTClaimString(claims, "team_id") } - return strings.TrimSpace(claims.Email) } diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go index d0caa5e527..54baef03a2 100644 --- a/backend/internal/service/grok_oauth_service_test.go +++ b/backend/internal/service/grok_oauth_service_test.go @@ -4,7 +4,10 @@ package service import ( "context" + "encoding/base64" + "encoding/json" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/stretchr/testify/require" @@ -12,6 +15,7 @@ import ( type grokOAuthClientStub struct { refreshResponse *xai.TokenResponse + ssoResponse *xai.TokenResponse exchangeCalls int } @@ -24,6 +28,10 @@ func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, stri return s.refreshResponse, nil } +func (s *grokOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { + return s.ssoResponse, nil +} + func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated(t *testing.T) { svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ refreshResponse: &xai.TokenResponse{ @@ -66,3 +74,43 @@ func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSessi require.Contains(t, err.Error(), "GROK_OAUTH_SESSION_NOT_FOUND") require.Zero(t, client.exchangeCalls) } + +func TestGrokOAuthServiceBuildAccountCredentialsDefaultsToSubscriptionProxy(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{}) + defer svc.Stop() + + credentials := svc.BuildAccountCredentials(&GrokTokenInfo{ + AccessToken: "access-token", + ExpiresAt: time.Now().Add(time.Hour).Unix(), + }) + + require.Equal(t, xai.DefaultCLIBaseURL, credentials["base_url"]) +} + +func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + ssoResponse: &xai.TokenResponse{ + AccessToken: makeGrokOAuthJWT(map[string]any{"sub": "user-sub", "team_id": "team-1"}), + RefreshToken: "refresh-token", + IDToken: makeGrokOAuthJWT(map[string]any{"email": "user@example.com"}), + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + info, err := svc.ConvertFromSSO(context.Background(), "sso-token", nil) + require.NoError(t, err) + require.Equal(t, "user@example.com", info.Email) + require.Equal(t, "user-sub", info.Subject) + require.Equal(t, "team-1", info.TeamID) + + credentials := svc.BuildAccountCredentials(info) + require.Equal(t, "user@example.com", credentials["email"]) + require.Equal(t, "user-sub", credentials["sub"]) + require.Equal(t, "team-1", credentials["team_id"]) +} + +func makeGrokOAuthJWT(claims map[string]any) string { + payload, _ := json.Marshal(claims) + return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".signature" +} diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go index 0939b78e20..f220fe33b9 100644 --- a/backend/internal/service/grok_quota_fetcher.go +++ b/backend/internal/service/grok_quota_fetcher.go @@ -3,6 +3,8 @@ package service import ( "encoding/json" "fmt" + "net/http" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -24,54 +26,150 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { } if account == nil { usage.ErrorCode = "quota_unknown" - usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers" return usage } + billing, _ := grokBillingSnapshotFromExtra(account.Extra) snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) + if billing != nil { + usage.GrokBilling = billing + if billing.Plan != "" { + usage.SubscriptionTier = billing.Plan + usage.SubscriptionTierRaw = billing.Plan + } + if parsedAt, parseErr := time.Parse(time.RFC3339, billing.UpdatedAt); parseErr == nil { + usage.UpdatedAt = &parsedAt + } + if billing.FetchedAt != "" { + usage.GrokLastQuotaProbeAt = billing.FetchedAt + } + usage.GrokQuotaSnapshotState = "billing_observed" + usage.GrokLastStatusCode = billing.StatusCode + switch billing.StatusCode { + case 401: + usage.NeedsReauth = true + usage.ErrorCode = "unauthenticated" + case 403: + usage.IsForbidden = true + usage.ForbiddenType = "forbidden" + usage.ErrorCode = "forbidden" + case 429: + usage.ErrorCode = "rate_limited" + } + } + if err != nil || snapshot == nil { - usage.ErrorCode = "quota_unknown" - usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + applyGrokCredentialUsageFallback(usage, account) + if billing == nil { + usage.ErrorCode = "quota_unknown" + usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers" + } return usage } - if parsedAt, err := time.Parse(time.RFC3339, snapshot.UpdatedAt); err == nil { - usage.UpdatedAt = &parsedAt + if parsedAt, parseErr := time.Parse(time.RFC3339, snapshot.UpdatedAt); parseErr == nil { + if billing == nil || usage.UpdatedAt == nil || parsedAt.After(*usage.UpdatedAt) { + usage.UpdatedAt = &parsedAt + } } usage.GrokRequestQuota = snapshot.Requests usage.GrokTokenQuota = snapshot.Tokens usage.GrokRetryAfterSeconds = snapshot.RetryAfterSeconds - usage.SubscriptionTier = snapshot.SubscriptionTier - usage.SubscriptionTierRaw = snapshot.SubscriptionTier - usage.GrokEntitlementStatus = snapshot.EntitlementStatus - usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt + if usage.SubscriptionTier == "" { + usage.SubscriptionTier = snapshot.SubscriptionTier + usage.SubscriptionTierRaw = snapshot.SubscriptionTier + } + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = snapshot.EntitlementStatus + } + if usage.GrokLastQuotaProbeAt == "" { + usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt + } usage.GrokLastHeadersSeenAt = snapshot.LastHeadersSeenAt - usage.GrokLastStatusCode = snapshot.StatusCode + if snapshot.StatusCode >= http.StatusBadRequest || usage.GrokLastStatusCode == 0 { + usage.GrokLastStatusCode = snapshot.StatusCode + } if snapshot.HasObservedHeaders() { - usage.GrokQuotaSnapshotState = "observed" - } else { + if usage.GrokQuotaSnapshotState == "" { + usage.GrokQuotaSnapshotState = "observed" + } + } else if billing == nil { usage.GrokQuotaSnapshotState = "no_headers" usage.ErrorCode = "quota_unknown" usage.Error = "No xAI quota headers observed on the latest Grok probe" } - switch snapshot.StatusCode { - case 401: - usage.NeedsReauth = true - usage.ErrorCode = "unauthenticated" - case 403: - usage.IsForbidden = true - usage.ForbiddenType = "forbidden" - usage.ErrorCode = "forbidden" - if usage.GrokEntitlementStatus == "" { - usage.GrokEntitlementStatus = "forbidden" + if usage.ErrorCode == "" { + switch snapshot.StatusCode { + case 401: + usage.NeedsReauth = true + usage.ErrorCode = "unauthenticated" + case 403: + usage.IsForbidden = true + usage.ForbiddenType = "forbidden" + usage.ErrorCode = "forbidden" + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = "forbidden" + } + case 429: + usage.ErrorCode = "rate_limited" } - case 429: - usage.ErrorCode = "rate_limited" } + applyGrokCredentialUsageFallback(usage, account) return usage } +func applyGrokCredentialUsageFallback(usage *UsageInfo, account *Account) { + if usage == nil || account == nil { + return + } + if usage.SubscriptionTier == "" { + tier := strings.TrimSpace(account.GetCredential("subscription_tier")) + usage.SubscriptionTier = tier + usage.SubscriptionTierRaw = tier + } + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = strings.TrimSpace(account.GetCredential("entitlement_status")) + } +} + +func grokBillingSnapshotFromExtra(extra map[string]any) (*xai.BillingSummary, error) { + if extra == nil { + return nil, nil + } + raw, ok := extra[grokBillingExtraKey] + if !ok || raw == nil { + return nil, nil + } + switch snapshot := raw.(type) { + case *xai.BillingSummary: + return snapshot, nil + case xai.BillingSummary: + return &snapshot, nil + case map[string]any: + data, err := json.Marshal(snapshot) + if err != nil { + return nil, err + } + var out xai.BillingSummary + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + default: + data, err := json.Marshal(raw) + if err != nil { + return nil, fmt.Errorf("marshal grok billing snapshot: %w", err) + } + var out xai.BillingSummary + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + } +} + func grokQuotaSnapshotFromExtra(extra map[string]any) (*xai.QuotaSnapshot, error) { if extra == nil { return nil, nil diff --git a/backend/internal/service/grok_quota_fetcher_test.go b/backend/internal/service/grok_quota_fetcher_test.go index d2d9c14993..1de9b51c9e 100644 --- a/backend/internal/service/grok_quota_fetcher_test.go +++ b/backend/internal/service/grok_quota_fetcher_test.go @@ -20,7 +20,34 @@ func TestGrokQuotaFetcherBuildUsageInfoUnknownUntilFirstSnapshot(t *testing.T) { usage := NewGrokQuotaFetcher().BuildUsageInfo(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth}) require.Equal(t, "passive", usage.Source) require.Equal(t, "quota_unknown", usage.ErrorCode) - require.Contains(t, usage.Error, "unknown until the first upstream response") + require.Contains(t, usage.Error, "unknown until billing is probed") +} + +func TestGrokQuotaFetcherUsesCredentialTierWhenBillingHasNoPlan(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "subscription_tier": " FREE ", + "entitlement_status": " active ", + }, + Extra: map[string]any{ + grokBillingExtraKey: &xai.BillingSummary{ + PeriodType: "weekly", + StatusCode: http.StatusOK, + UpdatedAt: "2030-01-01T00:00:00Z", + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + + require.NotNil(t, usage.GrokBilling) + require.Equal(t, "FREE", usage.SubscriptionTier) + require.Equal(t, "FREE", usage.SubscriptionTierRaw) + require.Equal(t, "active", usage.GrokEntitlementStatus) } func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) { @@ -68,6 +95,32 @@ func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) { require.True(t, usage.UpdatedAt.Equal(time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC))) } +func TestGrokQuotaFetcherSnapshotErrorOverridesSuccessfulBillingStatus(t *testing.T) { + t.Parallel() + + updatedAt := "2030-01-01T00:00:00Z" + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokBillingExtraKey: &xai.BillingSummary{ + PeriodType: "weekly", + StatusCode: http.StatusOK, + UpdatedAt: updatedAt, + }, + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + UpdatedAt: updatedAt, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + + require.Equal(t, "rate_limited", usage.ErrorCode) + require.Equal(t, http.StatusTooManyRequests, usage.GrokLastStatusCode) +} + func TestGrokQuotaFetcherBuildUsageInfoFromNoHeadersProbe(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 06d90a609e..2219a75c15 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -7,27 +7,37 @@ import ( "io" "log/slog" "net/http" + "strconv" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "golang.org/x/sync/singleflight" ) const ( grokQuotaUpstreamTimeout = 20 * time.Second grokQuotaProbeInput = "." - grokQuotaDefaultModel = "grok-4.3" + grokQuotaDefaultModel = grokDefaultResponsesModel + grokBillingExtraKey = "grok_billing_snapshot" ) type GrokQuotaProbeResult struct { - Source string `json:"source"` - Model string `json:"model"` - Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"` - StatusCode int `json:"status_code,omitempty"` - HeadersObserved bool `json:"headers_observed"` - ResetSupported bool `json:"reset_supported"` - FetchedAt int64 `json:"fetched_at"` + Source string `json:"source"` + Model string `json:"model,omitempty"` + Billing *xai.BillingSummary `json:"billing,omitempty"` + Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"` + LocalUsage24h *WindowStats `json:"local_usage_24h,omitempty"` + LocalUsage7d *WindowStats `json:"local_usage_7d,omitempty"` + LocalUsageMonthly *WindowStats `json:"local_usage_monthly,omitempty"` + StatusCode int `json:"status_code,omitempty"` + HeadersObserved bool `json:"headers_observed"` + ResetSupported bool `json:"reset_supported"` + FetchedAt int64 `json:"fetched_at"` + Persisted bool `json:"persisted"` + ProbeError string `json:"probe_error,omitempty"` } type GrokQuotaResetResult struct { @@ -41,6 +51,8 @@ type GrokQuotaService struct { proxyRepo ProxyRepository tokenProvider *GrokTokenProvider httpUpstream HTTPUpstream + usageLogRepo UsageLogRepository + probeFlight singleflight.Group } func NewGrokQuotaService( @@ -48,16 +60,71 @@ func NewGrokQuotaService( proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, + usageLogRepos ...UsageLogRepository, ) *GrokQuotaService { + var usageLogRepo UsageLogRepository + if len(usageLogRepos) > 0 { + usageLogRepo = usageLogRepos[0] + } return &GrokQuotaService{ accountRepo: accountRepo, proxyRepo: proxyRepo, tokenProvider: tokenProvider, httpUpstream: httpUpstream, + usageLogRepo: usageLogRepo, } } +// QueryQuota combines xAI billing data with an active quota-header probe for +// Free accounts, whose billing response does not include usage_percent. +func (s *GrokQuotaService) QueryQuota(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + billingResult, billingErr := s.ProbeBilling(ctx, accountID) + if billingErr == nil && billingResult != nil && grokBillingHasAuthoritativeQuota(billingResult.Billing) { + return billingResult, nil + } + + probeResult, probeErr := s.ProbeUsage(ctx, accountID) + if probeErr != nil { + if billingResult != nil && billingResult.Billing != nil { + billingResult.ProbeError = probeErr.Error() + return billingResult, nil + } + return nil, probeErr + } + if probeResult == nil { + if billingErr != nil { + return nil, billingErr + } + return nil, infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_EMPTY", "Grok quota probe returned no result") + } + if billingResult != nil { + probeResult.Source = "hybrid_probe" + probeResult.Billing = billingResult.Billing + probeResult.LocalUsage24h = billingResult.LocalUsage24h + probeResult.LocalUsage7d = billingResult.LocalUsage7d + probeResult.LocalUsageMonthly = billingResult.LocalUsageMonthly + probeResult.Persisted = probeResult.Persisted || billingResult.Persisted + } + return probeResult, nil +} + +func grokBillingHasAuthoritativeQuota(billing *xai.BillingSummary) bool { + if billing == nil { + return false + } + return billing.UsagePercent != nil || + billing.UsedPercent != nil || + (billing.MonthlyLimitCents != nil && *billing.MonthlyLimitCents > 0) || + strings.TrimSpace(billing.Plan) != "" +} + func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + return s.runProbeFlight(ctx, "active:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) { + return s.probeUsage(sharedCtx, accountID) + }) +} + +func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { account, token, proxyURL, err := s.prepareProbe(ctx, accountID) if err != nil { return nil, err @@ -82,7 +149,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr req.Header.Set("Authorization", "Bearer "+token) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") - req.Header.Set("User-Agent", "sub2api-grok-quota-probe/1.0") + applyGrokCLIHeaders(req.Header) resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1)) if err != nil { @@ -91,9 +158,16 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr defer func() { _ = resp.Body.Close() }() snapshot := xai.ObserveQuotaHeaders(resp.Header, resp.StatusCode, "active_probe") - _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + resetAt, limited := grokRateLimitResetAt(snapshot, time.Now()) + if limited { + normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now()) + } + persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ grokQuotaSnapshotExtraKey: snapshot, }) + if limited { + persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + } result := &GrokQuotaProbeResult{ Source: "active_probe", @@ -103,19 +177,201 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr HeadersObserved: snapshot.HeadersObserved, ResetSupported: false, FetchedAt: time.Now().Unix(), + Persisted: persistErr == nil, } if resp.StatusCode == http.StatusTooManyRequests { return result, nil } if resp.StatusCode >= 400 { - bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, 240)) - bodyText := truncate(strings.TrimSpace(string(bodyBytes)), 240) - slog.Warn("grok_quota_probe_failed", "account_id", account.ID, "model", probeModel, "status", resp.StatusCode, "body", bodyText) - return nil, infraerrors.Newf(mapUpstreamStatus(resp.StatusCode), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "upstream returned %d for probe model %q: %s", resp.StatusCode, probeModel, bodyText) + const reason = "GROK_QUOTA_PROBE_UPSTREAM_ERROR" + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 4<<10)) + slog.Warn( + "grok_quota_probe_failed", + "account_id", account.ID, + "model", probeModel, + "status", resp.StatusCode, + "reason", reason, + ) + return nil, infraerrors.Newf( + mapUpstreamStatus(resp.StatusCode), + reason, + "upstream returned %d for probe model %q", + resp.StatusCode, + probeModel, + ) } return result, nil } +// ProbeBilling only calls the xAI billing endpoints. Account usage refreshes +// use this method so opening the account list never consumes model quota. +func (s *GrokQuotaService) ProbeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + return s.runProbeFlight(ctx, "billing:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) { + return s.probeBilling(sharedCtx, accountID) + }) +} + +func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + account, token, proxyURL, err := s.prepareProbe(ctx, accountID) + if err != nil { + return nil, err + } + + probeCtx, cancel := context.WithTimeout(ctx, grokQuotaUpstreamTimeout) + defer cancel() + type billingResult struct { + summary *xai.BillingSummary + status int + err error + } + var weekly, monthly billingResult + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + weekly.summary, weekly.status, weekly.err = s.fetchBilling(probeCtx, account, token, proxyURL, true) + }() + go func() { + defer wg.Done() + monthly.summary, monthly.status, monthly.err = s.fetchBilling(probeCtx, account, token, proxyURL, false) + }() + wg.Wait() + + weeklyOK := weekly.summary != nil + monthlyOK := monthly.summary != nil + if !weeklyOK && !monthlyOK { + return nil, mergeGrokBillingProbeErrors(weekly.status, monthly.status, weekly.err, monthly.err) + } + statusCode := preferSuccessfulBillingStatus(weekly.status, monthly.status, weeklyOK, monthlyOK) + previous, _ := grokBillingSnapshotFromExtra(account.Extra) + billing := xai.MergeBillingProbeResult(previous, weekly.summary, monthly.summary, weeklyOK, monthlyOK) + billing = xai.StampBillingSummary(billing, statusCode, "billing_probe") + persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + grokBillingExtraKey: billing, + }) + if persistErr != nil { + slog.Warn("grok_billing_persist_failed", "account_id", account.ID, "error", persistErr) + } + now := time.Now().UTC() + localUsage24h, localUsage7d, localUsageMonthly := grokLocalUsageForQuota(ctx, s.usageLogRepo, account.ID, billing, now) + return &GrokQuotaProbeResult{ + Source: "billing_probe", + Billing: billing, + LocalUsage24h: localUsage24h, + LocalUsage7d: localUsage7d, + LocalUsageMonthly: localUsageMonthly, + StatusCode: statusCode, + FetchedAt: now.Unix(), + Persisted: persistErr == nil, + }, nil +} + +func (s *GrokQuotaService) runProbeFlight( + ctx context.Context, + key string, + probe func(context.Context) (*GrokQuotaProbeResult, error), +) (*GrokQuotaProbeResult, error) { + if s == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_NOT_CONFIGURED", "grok quota service is not configured") + } + resultCh := s.probeFlight.DoChan(key, func() (any, error) { + sharedCtx, cancel := context.WithTimeout(context.Background(), grokQuotaUpstreamTimeout+5*time.Second) + defer cancel() + return probe(sharedCtx) + }) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case flightResult := <-resultCh: + if flightResult.Err != nil { + return nil, flightResult.Err + } + result, ok := flightResult.Val.(*GrokQuotaProbeResult) + if !ok || result == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_PROBE_RESULT_INVALID", "invalid Grok quota probe result") + } + cloned := *result + return &cloned, nil + } +} + +func (s *GrokQuotaService) fetchBilling( + ctx context.Context, + account *Account, + token string, + proxyURL string, + weekly bool, +) (*xai.BillingSummary, int, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, xai.BuildBillingURL(weekly), nil) + if err != nil { + return nil, 0, infraerrors.Newf(http.StatusInternalServerError, "GROK_QUOTA_PROBE_REQUEST_BUILD_FAILED", "failed to build billing request: %v", err) + } + xai.ApplyCLIBillingHeaders(req, token) + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 2)) + if err != nil { + return nil, 0, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_PROBE_REQUEST_FAILED", "billing request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + + bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if resp.StatusCode == http.StatusTooManyRequests { + return nil, resp.StatusCode, nil + } + if resp.StatusCode >= 400 { + bodyText := truncate(strings.TrimSpace(string(bodyBytes)), 240) + slog.Warn("grok_quota_billing_failed", "account_id", account.ID, "weekly", weekly, "status", resp.StatusCode, "body", bodyText) + return nil, resp.StatusCode, infraerrors.Newf(mapUpstreamStatus(resp.StatusCode), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing returned %d: %s", resp.StatusCode, bodyText) + } + payload, err := xai.ParseBillingPayload(bodyBytes) + if err != nil { + return nil, resp.StatusCode, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_BILLING_PARSE_ERROR", "failed to parse billing body: %v", err) + } + return xai.BuildBillingSummary(payload.Config), resp.StatusCode, nil +} + +func mergeGrokBillingProbeErrors(weeklyStatus, monthlyStatus int, weeklyErr, monthlyErr error) error { + weeklyKey := grokBillingProbeErrorKey(weeklyStatus, weeklyErr) + monthlyKey := grokBillingProbeErrorKey(monthlyStatus, monthlyErr) + if weeklyKey == monthlyKey { + switch { + case weeklyErr != nil: + return weeklyErr + case monthlyErr != nil: + return monthlyErr + case weeklyStatus == http.StatusTooManyRequests: + return infraerrors.New(http.StatusTooManyRequests, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing rate limited") + case weeklyStatus != 0 && weeklyStatus != http.StatusOK: + return infraerrors.New(mapUpstreamStatus(weeklyStatus), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "xAI billing endpoints returned the same upstream error") + default: + return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_BILLING_EMPTY", "xAI billing endpoints returned no quota data") + } + } + slog.Warn("grok_quota_probe_parts_failed", "weekly_status", weeklyStatus, "weekly_error", weeklyErr, "monthly_status", monthlyStatus, "monthly_error", monthlyErr) + return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_PARTS_FAILED", "weekly and monthly billing probes failed differently").WithMetadata(map[string]string{ + "weekly_status": strconv.Itoa(weeklyStatus), "monthly_status": strconv.Itoa(monthlyStatus), + }) +} + +func grokBillingProbeErrorKey(status int, err error) string { + if err != nil { + return strconv.Itoa(status) + ":" + strconv.Itoa(infraerrors.Code(err)) + ":" + infraerrors.Reason(err) + } + return strconv.Itoa(status) + ":empty" +} + +func preferSuccessfulBillingStatus(weeklyStatus, monthlyStatus int, weeklyOK, monthlyOK bool) int { + if weeklyOK && weeklyStatus >= 200 && weeklyStatus < 300 { + return weeklyStatus + } + if monthlyOK && monthlyStatus >= 200 && monthlyStatus < 300 { + return monthlyStatus + } + if weeklyStatus != 0 { + return weeklyStatus + } + return monthlyStatus +} + func (s *GrokQuotaService) ResetQuota(ctx context.Context, accountID int64) (*GrokQuotaResetResult, error) { if _, err := s.loadGrokOAuthAccount(ctx, accountID); err != nil { return nil, err diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index fe1a00aa89..ba9b6cbcb8 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -3,14 +3,19 @@ package service import ( + "bytes" "context" "io" + "log/slog" "net/http" + "strconv" "strings" + "sync" "testing" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -19,6 +24,10 @@ import ( type grokQuotaAccountRepo struct { *mockAccountRepoForPlatform updates map[int64]map[string]any + updateCalls int + rateLimitedCalls int + lastRateLimitedID int64 + lastRateLimitResetAt time.Time tempUnschedCalls int lastTempUnschedID int64 lastTempUnschedUntil time.Time @@ -26,6 +35,7 @@ type grokQuotaAccountRepo struct { } func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + r.updateCalls++ if r.updates == nil { r.updates = make(map[int64]map[string]any) } @@ -33,6 +43,17 @@ func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates return nil } +func (r *grokQuotaAccountRepo) SetRateLimited(_ context.Context, id int64, resetAt time.Time) error { + r.rateLimitedCalls++ + r.lastRateLimitedID = id + r.lastRateLimitResetAt = resetAt + return nil +} + +func (r *grokQuotaAccountRepo) SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error { + return r.SetRateLimited(ctx, id, resetAt) +} + func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, until time.Time, reason string) error { r.tempUnschedCalls++ r.lastTempUnschedID = id @@ -47,6 +68,113 @@ type grokQuotaProxyRepo struct { calls int } +type grokQuotaUsageLogRepo struct { + UsageLogRepository + stats *usagestats.AccountStats + err error + calls int + startTimes []time.Time +} + +func (r *grokQuotaUsageLogRepo) GetAccountWindowStats(_ context.Context, _ int64, start time.Time) (*usagestats.AccountStats, error) { + r.calls++ + r.startTimes = append(r.startTimes, start) + return r.stats, r.err +} + +func (r *grokQuotaUsageLogRepo) GetAccountTodayStats(context.Context, int64) (*usagestats.AccountStats, error) { + return nil, nil +} + +type grokHybridUpstream struct { + httpUpstreamRecorder + mu sync.Mutex + requests []*http.Request + bodies [][]byte + weeklyUsagePercent *float64 + monthlyLimitCents *float64 + activeStatus int + activeHeaders http.Header + billingStarted chan struct{} + billingRelease <-chan struct{} + billingStartOnce sync.Once + billingStatus int + billingHeaders http.Header +} + +func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + var body []byte + if req != nil && req.Body != nil { + body, _ = io.ReadAll(req.Body) + } + u.mu.Lock() + u.requests = append(u.requests, req) + u.bodies = append(u.bodies, body) + u.mu.Unlock() + + if req.URL.Path == "/v1/responses" { + status := u.activeStatus + if status == 0 { + status = http.StatusOK + } + headers := u.activeHeaders + if headers == nil { + headers = http.Header{ + "X-Ratelimit-Limit-Tokens": []string{"2000000"}, + "X-Ratelimit-Remaining-Tokens": []string{"1500000"}, + } + } + return &http.Response{StatusCode: status, Header: headers, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`))}, nil + } + if u.billingStarted != nil { + u.billingStartOnce.Do(func() { close(u.billingStarted) }) + } + if u.billingRelease != nil { + select { + case <-u.billingRelease: + case <-req.Context().Done(): + return nil, req.Context().Err() + } + } + if u.billingStatus != 0 && u.billingStatus != http.StatusOK { + return &http.Response{ + StatusCode: u.billingStatus, + Header: u.billingHeaders, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"billing limited"}}`)), + }, nil + } + + if req.URL.RawQuery == "format=credits" { + usage := "" + if u.weeklyUsagePercent != nil { + usage = `,"creditUsagePercent":` + strconv.FormatFloat(*u.weeklyUsagePercent, 'f', -1, 64) + } + payload := `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}` + usage + `}}` + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil + } + monthlyLimit := "" + if u.monthlyLimitCents != nil { + monthlyLimit = `,"monthlyLimit":{"val":` + strconv.FormatFloat(*u.monthlyLimitCents, 'f', -1, 64) + `}` + } + monthlyPayload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"` + monthlyLimit + `}}` + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(monthlyPayload)), + }, nil +} + +func (u *grokHybridUpstream) snapshot() ([]*http.Request, [][]byte) { + u.mu.Lock() + defer u.mu.Unlock() + requests := append([]*http.Request(nil), u.requests...) + bodies := make([][]byte, len(u.bodies)) + for i := range u.bodies { + bodies[i] = append([]byte(nil), u.bodies[i]...) + } + return requests, bodies +} + func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error) { r.calls++ return r.proxies[id], nil @@ -86,7 +214,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { result, err := svc.ProbeUsage(context.Background(), 42) require.NoError(t, err) require.Equal(t, http.StatusOK, result.StatusCode) - require.Equal(t, "grok-4.3", result.Model) + require.Equal(t, "grok-4.5", result.Model) require.True(t, result.HeadersObserved) require.NotNil(t, result.Snapshot) require.True(t, result.Snapshot.HeadersObserved) @@ -96,9 +224,10 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { require.NotNil(t, result.Snapshot.Requests) require.EqualValues(t, 10, *result.Snapshot.Requests.Limit) require.EqualValues(t, 7, *result.Snapshot.Requests.Remaining) - require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`) require.Contains(t, string(upstream.lastBody), `"store":false`) require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey]) @@ -135,8 +264,8 @@ func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) { result, err := svc.ProbeUsage(context.Background(), 47) require.NoError(t, err) - require.Equal(t, "grok-4.3", result.Model) - require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "grok-4.5", result.Model) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.NotContains(t, string(upstream.lastBody), "grok-composer") } @@ -168,7 +297,55 @@ func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T _, err := svc.ProbeUsage(context.Background(), 48) require.Error(t, err) require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err)) - require.Contains(t, infraerrors.Message(err), `probe model "grok-4.3"`) + require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`) +} + +func TestGrokQuotaServiceProbeUsageRedactsUpstreamErrorBodyFromErrorAndLogs(t *testing.T) { + const upstreamSecret = "upstream-secret-refresh-token" + account := &Account{ + ID: 49, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{49: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{}, + Body: io.NopCloser(strings.NewReader( + `{"error":"` + upstreamSecret + `","detail":"credential rejected"}`, + )), + }} + svc := NewGrokQuotaService( + repo, + nil, + NewGrokTokenProvider(repo, nil), + upstream, + ) + + var logs bytes.Buffer + previousLogger := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + defer slog.SetDefault(previousLogger) + + _, err := svc.ProbeUsage(context.Background(), account.ID) + require.Error(t, err) + require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err)) + require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`) + require.NotContains(t, err.Error(), upstreamSecret) + require.NotContains(t, infraerrors.Message(err), upstreamSecret) + require.Contains(t, logs.String(), "GROK_QUOTA_PROBE_UPSTREAM_ERROR") + require.NotContains(t, logs.String(), upstreamSecret) + require.NotContains(t, logs.String(), "credential rejected") + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) } func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T) { @@ -285,6 +462,389 @@ func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { require.NotNil(t, result.Snapshot) require.NotNil(t, result.Snapshot.RetryAfterSeconds) require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) +} + +func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 51, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{} + usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_000_000}} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, usageRepo) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "hybrid_probe", result.Source) + require.Equal(t, "grok-4.5", result.Model) + require.NotNil(t, result.Billing) + require.Nil(t, result.Billing.UsagePercent) + require.NotNil(t, result.LocalUsage24h) + require.EqualValues(t, 1_000_000, result.LocalUsage24h.Tokens) + require.Equal(t, 1, usageRepo.calls) + require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), usageRepo.startTimes[0], time.Second) + require.NotNil(t, result.Snapshot) + require.NotNil(t, result.Snapshot.Tokens) + require.EqualValues(t, 2_000_000, *result.Snapshot.Tokens.Limit) + require.True(t, result.HeadersObserved) + + requests, bodies := upstream.snapshot() + require.Len(t, requests, 3) + responseCalls := 0 + for i, req := range requests { + if req.URL.Path != "/v1/responses" { + continue + } + responseCalls++ + require.Equal(t, http.MethodPost, req.Method) + require.Equal(t, "grok-4.5", gjson.GetBytes(bodies[i], "model").String()) + require.EqualValues(t, 1, gjson.GetBytes(bodies[i], "max_output_tokens").Int()) + } + require.Equal(t, 1, responseCalls) +} + +func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 52, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + usagePercent := 25.0 + upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent} + usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_000_000}} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, usageRepo) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "billing_probe", result.Source) + require.NotNil(t, result.Billing) + require.InDelta(t, usagePercent, *result.Billing.UsagePercent, 1e-9) + require.Nil(t, result.Snapshot) + require.Empty(t, result.Model) + require.Nil(t, result.LocalUsage24h) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokQuotaServiceQueryQuotaCustomPaidMonthlyLimitSkipsActiveProbe(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 57, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + monthlyLimit := 25_000.0 + upstream := &grokHybridUpstream{monthlyLimitCents: &monthlyLimit} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "billing_probe", result.Source) + require.NotNil(t, result.Billing) + require.InDelta(t, monthlyLimit, *result.Billing.MonthlyLimitCents, 1e-9) + require.Nil(t, result.Snapshot) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokLocalUsage24hUsesRollingUTCWindow(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 7, 14, 20, 30, 0, 0, time.FixedZone("UTC+8", 8*60*60)) + + t.Run("returns usage from exact rolling window", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_250_000}} + stats := grokLocalUsage24h(context.Background(), repo, 57, now) + + require.NotNil(t, stats) + require.EqualValues(t, 1_250_000, stats.Tokens) + require.Equal(t, []time.Time{now.UTC().Add(-24 * time.Hour)}, repo.startTimes) + }) + + t.Run("query failure returns no stats", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded} + stats := grokLocalUsage24h(context.Background(), repo, 57, now) + + require.Nil(t, stats) + require.Equal(t, []time.Time{now.UTC().Add(-24 * time.Hour)}, repo.startTimes) + }) + + t.Run("missing repository returns no stats", func(t *testing.T) { + require.Nil(t, grokLocalUsage24h(context.Background(), nil, 57, now)) + }) + + t.Run("invalid account returns no stats without query", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{} + require.Nil(t, grokLocalUsage24h(context.Background(), repo, 0, now)) + require.Zero(t, repo.calls) + }) +} + +func TestGrokLocalUsageForQuotaSelectsFreeOrPaidWindows(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 7, 14, 12, 0, 0, 0, time.UTC) + billing := &xai.BillingSummary{ + PeriodType: "weekly", + PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339), + PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339), + BillingPeriodStart: now.Add(-13 * 24 * time.Hour).Format(time.RFC3339), + BillingPeriodEnd: now.Add(17 * 24 * time.Hour).Format(time.RFC3339), + } + + t.Run("free queries only rolling 24h", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 500_000}} + rolling, weekly, monthly := grokLocalUsageForQuota(context.Background(), repo, 57, billing, now) + + require.NotNil(t, rolling) + require.Nil(t, weekly) + require.Nil(t, monthly) + require.Equal(t, []time.Time{now.Add(-24 * time.Hour)}, repo.startTimes) + }) + + t.Run("paid queries only billing windows", func(t *testing.T) { + usagePercent := 25.0 + paidBilling := *billing + paidBilling.UsagePercent = &usagePercent + repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 500_000}} + rolling, weekly, monthly := grokLocalUsageForQuota(context.Background(), repo, 57, &paidBilling, now) + + require.Nil(t, rolling) + require.NotNil(t, weekly) + require.NotNil(t, monthly) + require.Equal(t, []time.Time{ + now.Add(-4 * 24 * time.Hour), + now.Add(-13 * 24 * time.Hour), + }, repo.startTimes) + }) +} + +func TestGrokLocalUsageForBillingOnlyReturnsAvailableWindows(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC) + billing := &xai.BillingSummary{ + PeriodType: "weekly", + PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339), + PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339), + } + + t.Run("valid weekly window", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_500_000}} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now) + require.NotNil(t, weekly) + require.EqualValues(t, 1_500_000, weekly.Tokens) + require.Nil(t, monthly) + require.Equal(t, 1, repo.calls) + }) + + t.Run("query failure", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now) + require.Nil(t, weekly) + require.Nil(t, monthly) + require.Equal(t, 1, repo.calls) + }) + + t.Run("missing billing window", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, nil, now) + require.Nil(t, weekly) + require.Nil(t, monthly) + require.Zero(t, repo.calls) + }) +} + +func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 54, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{} + usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 750_000}} + quotaService := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, usageRepo) + usageService := &AccountUsageService{ + grokQuotaFetcher: NewGrokQuotaFetcher(), + grokQuotaService: quotaService, + usageLogRepo: usageRepo, + cache: NewUsageCache(), + } + + usage, err := usageService.getGrokUsage(context.Background(), account, false) + require.NoError(t, err) + require.NotNil(t, usage.GrokBilling) + require.Nil(t, usage.GrokBilling.UsagePercent) + require.NotNil(t, usage.GrokLocalUsage24h) + require.EqualValues(t, 750_000, usage.GrokLocalUsage24h.Tokens) + require.Equal(t, 1, usageRepo.calls) + require.Len(t, usageRepo.startTimes, 1) + require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), usageRepo.startTimes[0], time.Second) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, http.MethodGet, req.Method) + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 55, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + billingStarted := make(chan struct{}) + billingRelease := make(chan struct{}) + upstream := &grokHybridUpstream{billingStarted: billingStarted, billingRelease: billingRelease} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + type probeOutcome struct { + result *GrokQuotaProbeResult + err error + } + billingOutcomes := make(chan probeOutcome, 2) + go func() { + result, err := svc.ProbeBilling(context.Background(), account.ID) + billingOutcomes <- probeOutcome{result: result, err: err} + }() + <-billingStarted + secondStarted := make(chan struct{}) + go func() { + close(secondStarted) + result, err := svc.ProbeBilling(context.Background(), account.ID) + billingOutcomes <- probeOutcome{result: result, err: err} + }() + <-secondStarted + time.Sleep(25 * time.Millisecond) + + activeResult, err := svc.ProbeUsage(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, activeResult.Snapshot) + close(billingRelease) + for range 2 { + outcome := <-billingOutcomes + require.NoError(t, outcome.err) + require.NotNil(t, outcome.result.Billing) + } + + requests, _ := upstream.snapshot() + billingCalls := 0 + activeCalls := 0 + for _, req := range requests { + switch req.URL.Path { + case "/v1/billing": + billingCalls++ + case "/v1/responses": + activeCalls++ + } + } + require.Equal(t, 2, billingCalls) + require.Equal(t, 1, activeCalls) +} + +func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 56, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{ + billingStatus: http.StatusTooManyRequests, + billingHeaders: http.Header{"Retry-After": []string{"45"}}, + } + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.ProbeBilling(context.Background(), account.ID) + + require.Error(t, err) + require.Nil(t, result) + require.Zero(t, repo.rateLimitedCalls) +} + +func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 53, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{ + activeStatus: http.StatusTooManyRequests, + activeHeaders: http.Header{"Retry-After": []string{"45"}}, + } + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, http.StatusTooManyRequests, result.StatusCode) + require.NotNil(t, result.Billing) + require.NotNil(t, result.Snapshot) + require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) } func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) { diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index a61e356a01..346eef7b95 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -50,6 +50,9 @@ type Group struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用); + // nil 表示使用默认价 defaultWebSearchPricePerCall(官方 $10/1000 次)。 + WebSearchPricePerCall *float64 // Claude Code 客户端限制 ClaudeCodeOnly bool diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go index 5a063a7d71..613f11d048 100644 --- a/backend/internal/service/image_generation_intent.go +++ b/backend/internal/service/image_generation_intent.go @@ -9,9 +9,23 @@ import ( const ( openAIResponsesEndpoint = "/v1/responses" openAIResponsesCompactEndpoint = "/v1/responses/compact" + responsesLiteHeader = "X-OpenAI-Internal-Codex-Responses-Lite" + responsesLiteHeaderKey = "x-openai-internal-codex-responses-lite" + responsesLiteWSMetadataKey = "ws_request_header_x_openai_internal_codex_responses_lite" imageGenerationPermissionMessage = "Image generation is not enabled for this group" ) +func isOpenAIResponsesLiteHeader(value string) bool { + return strings.EqualFold(strings.TrimSpace(value), "true") +} + +func isOpenAIResponsesLiteWebSocketPayload(body []byte) bool { + if len(body) == 0 || !gjson.ValidBytes(body) { + return false + } + return isOpenAIResponsesLiteHeader(gjson.GetBytes(body, "client_metadata."+responsesLiteWSMetadataKey).String()) +} + // ImageGenerationPermissionMessage returns the stable end-user error text for disabled groups. func ImageGenerationPermissionMessage() string { return imageGenerationPermissionMessage diff --git a/backend/internal/service/media_price_config.go b/backend/internal/service/media_price_config.go index ed84998906..0a583b6c72 100644 --- a/backend/internal/service/media_price_config.go +++ b/backend/internal/service/media_price_config.go @@ -29,3 +29,10 @@ func videoPriceConfigFromAPIKey(apiKey *APIKey) *VideoPriceConfig { func apiKeyHasConfiguredVideoPrice(apiKey *APIKey, resolution string) bool { return apiKey != nil && apiKey.Group != nil && apiKey.Group.GetVideoPrice(resolution) != nil } + +func webSearchPricePerCallFromAPIKey(apiKey *APIKey) *float64 { + if apiKey == nil || apiKey.Group == nil { + return nil + } + return apiKey.Group.WebSearchPricePerCall +} diff --git a/backend/internal/service/model_not_found_error.go b/backend/internal/service/model_not_found_error.go index 910a97d844..de4a004d1e 100644 --- a/backend/internal/service/model_not_found_error.go +++ b/backend/internal/service/model_not_found_error.go @@ -22,6 +22,30 @@ func isModelNotFoundError(statusCode int, body []byte) bool { return isUpstreamModelNotFoundError(statusCode, body) || statusCode == http.StatusNotFound } +// openAICodexPlanGatedModelPhrase matches the deterministic Codex 400 returned +// when a ChatGPT OAuth account's plan cannot serve the requested model, e.g. +// {"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."} +// The phrase is compared against the normalized body (lowercased, "_"/"-" +// folded to spaces), so it also matches the same message embedded in +// error.message-style payloads. +const openAICodexPlanGatedModelPhrase = "model is not supported when using codex" + +// isOpenAICodexPlanGatedModelError reports whether the upstream response is the +// deterministic Codex rejection of a plan-gated model on a ChatGPT account. +// Unlike transient failures, retrying the same account cannot succeed until the +// account's plan changes, so callers should treat it like model-not-found and +// cool the (account, model) pair down instead of re-selecting the account. +func isOpenAICodexPlanGatedModelError(statusCode int, body []byte) bool { + if statusCode != http.StatusBadRequest { + return false + } + normalized := normalizeModelNotFoundBody(body) + if normalized == "" { + return false + } + return strings.Contains(normalized, openAICodexPlanGatedModelPhrase) +} + func containsModelNotFoundKeyword(normalizedBody string) bool { if normalizedBody == "" { return false diff --git a/backend/internal/service/model_not_found_error_test.go b/backend/internal/service/model_not_found_error_test.go index a87340eb55..2f8c83466f 100644 --- a/backend/internal/service/model_not_found_error_test.go +++ b/backend/internal/service/model_not_found_error_test.go @@ -64,3 +64,51 @@ func TestAntigravityModelNotFoundKeepsBare404Fallback(t *testing.T) { t.Fatal("antigravity model-not-found helper should keep bare 404 fallback") } } + +func TestIsOpenAICodexPlanGatedModelError(t *testing.T) { + tests := []struct { + name string + statusCode int + body []byte + want bool + }{ + { + name: "400 codex plan gated detail payload", + statusCode: http.StatusBadRequest, + body: []byte(`{"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."}`), + want: true, + }, + { + name: "400 codex plan gated error message payload", + statusCode: http.StatusBadRequest, + body: []byte(`{"error":{"message":"The 'gpt-5.4' model is not supported when using Codex with a ChatGPT account."}}`), + want: true, + }, + { + name: "400 unrelated invalid request does not match", + statusCode: http.StatusBadRequest, + body: []byte(`{"error":{"message":"Invalid schema for response_format 'agentic_plan'"}}`), + want: false, + }, + { + name: "404 with plan gated message does not match", + statusCode: http.StatusNotFound, + body: []byte(`{"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."}`), + want: false, + }, + { + name: "400 empty body does not match", + statusCode: http.StatusBadRequest, + body: nil, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isOpenAICodexPlanGatedModelError(tt.statusCode, tt.body); got != tt.want { + t.Fatalf("isOpenAICodexPlanGatedModelError() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/backend/internal/service/oauth_service.go b/backend/internal/service/oauth_service.go index 1369dd9e89..0b3888a73f 100644 --- a/backend/internal/service/oauth_service.go +++ b/backend/internal/service/oauth_service.go @@ -22,6 +22,7 @@ type OpenAIOAuthClient interface { type GrokOAuthClient interface { ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error) + ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) } // GrokOAuthTokenService is the narrow refresh port used by Grok token providers. diff --git a/backend/internal/service/openai_agent_identity.go b/backend/internal/service/openai_agent_identity.go index 4bf447ad64..f375eac85c 100644 --- a/backend/internal/service/openai_agent_identity.go +++ b/backend/internal/service/openai_agent_identity.go @@ -492,16 +492,20 @@ func redactAgentIdentitySensitiveBodyForAccount(ctx context.Context, repo Accoun redacted = strings.ReplaceAll(redacted, value, "[redacted]") } } - for { - start := strings.Index(redacted, "AgentAssertion ") - if start < 0 { + const assertionPrefix = "AgentAssertion " + for offset := 0; offset < len(redacted); { + relativeStart := strings.Index(redacted[offset:], assertionPrefix) + if relativeStart < 0 { break } - end := start + len("AgentAssertion ") + start := offset + relativeStart + valueStart := start + len(assertionPrefix) + end := valueStart for end < len(redacted) && !strings.ContainsRune(" \t\r\n\"',}", rune(redacted[end])) { end++ } - redacted = redacted[:start] + "AgentAssertion [redacted]" + redacted[end:] + redacted = redacted[:valueStart] + "[redacted]" + redacted[end:] + offset = valueStart + len("[redacted]") } return []byte(redacted) } diff --git a/backend/internal/service/openai_agent_identity_compat_test.go b/backend/internal/service/openai_agent_identity_compat_test.go index eb90dfb56e..5dad35887b 100644 --- a/backend/internal/service/openai_agent_identity_compat_test.go +++ b/backend/internal/service/openai_agent_identity_compat_test.go @@ -133,7 +133,32 @@ func TestOpenAIAgentIdentityPassthroughKeepsSessionAndPromptCacheHeaders(t *test require.Equal(t, "account-agent-passthrough", req.Header.Get("chatgpt-account-id")) require.NotEqual(t, "client-session", req.Header.Get("session_id")) require.NotEqual(t, "client-conversation", req.Header.Get("conversation_id")) - require.Equal(t, isolateOpenAISessionID(0, "cache-agent"), req.Header.Get("session_id")) + require.Equal(t, isolateOpenAISessionID(0, "client-session"), req.Header.Get("session_id")) + require.Equal(t, isolateOpenAISessionID(0, "client-conversation"), req.Header.Get("conversation_id")) + requestBody, err := io.ReadAll(req.Body) + require.NoError(t, err) + require.Contains(t, string(requestBody), `"prompt_cache_key":"cache-agent"`) + + // Authentication mode must not affect session isolation or prompt-cache + // behavior. Compare the same request with the existing OAuth path instead + // of pinning this test to an implementation-specific hash. + oauthAccount := &Account{ + ID: 26, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "account-oauth-passthrough", + }, + } + oauthRecorder := httptest.NewRecorder() + oauthContext, _ := gin.CreateTestContext(oauthRecorder) + oauthContext.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + oauthContext.Request.Header.Set("session_id", "client-session") + oauthContext.Request.Header.Set("conversation_id", "client-conversation") + oauthReq, err := svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), oauthContext, oauthAccount, body, "oauth-token") + require.NoError(t, err) + require.Equal(t, oauthReq.Header.Get("session_id"), req.Header.Get("session_id")) + require.Equal(t, oauthReq.Header.Get("conversation_id"), req.Header.Get("conversation_id")) } func TestOpenAIAgentIdentityErrorRedactionDoesNotLeakCredentialValues(t *testing.T) { diff --git a/backend/internal/service/openai_alpha_search.go b/backend/internal/service/openai_alpha_search.go new file mode 100644 index 0000000000..50d37f95bb --- /dev/null +++ b/backend/internal/service/openai_alpha_search.go @@ -0,0 +1,165 @@ +package service + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" +) + +const ( + chatgptCodexAlphaSearchURL = "https://chatgpt.com/backend-api/codex/alpha/search" + openAIPlatformAlphaSearchURL = "https://api.openai.com/v1/alpha/search" +) + +// ForwardAlphaSearch proxies Codex standalone web search without binding the +// evolving alpha request or response schema. +// +// 返回值约定:仅当上游返回 2xx(一次真实成功的搜索)时返回非 nil 的 +// *OpenAIForwardResult(WebSearchCalls=1,供按次计费);上游错误被原样透传 +// 给客户端时返回 (nil, nil),不产生计费。 +func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { + if s == nil || c == nil || account == nil { + return nil, fmt.Errorf("service, context, and account are required") + } + modelResult := gjson.GetBytes(body, "model") + requestedModel := strings.TrimSpace(modelResult.String()) + if modelResult.Type != gjson.String || requestedModel == "" { + return nil, fmt.Errorf("model is required") + } + + upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(requestedModel)) + if upstreamModel != "" && upstreamModel != requestedModel { + body = ReplaceModelInBody(body, upstreamModel) + } + + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + + req, err := s.buildOpenAIAlphaSearchRequest(ctx, c, account, body, token) + if err != nil { + return nil, err + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + upstreamStart := time.Now() + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + return nil, fmt.Errorf("read alpha search response: %w", err) + } + + if resp.StatusCode >= http.StatusBadRequest { + upstreamMessage := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) + if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMessage, respBody) { + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + s.handleFailoverSideEffects(ctx, resp, account, respBody, upstreamModel) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + } + + if !account.IsShadow() { + s.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, resp.Header) + } + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + contentType := resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "application/json" + } + c.Data(resp.StatusCode, contentType, respBody) + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + // 非 2xx(错误/重定向)已原样透传给客户端:不是一次成功的搜索,不计费。 + return nil, nil + } + return &OpenAIForwardResult{ + RequestID: strings.TrimSpace(resp.Header.Get("x-request-id")), + Model: requestedModel, + UpstreamModel: upstreamModel, + Duration: time.Since(upstreamStart), + WebSearchCalls: 1, + }, nil +} + +func (s *OpenAIGatewayService) buildOpenAIAlphaSearchRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { + clientBeta := "" + if c != nil { + clientBeta = c.GetHeader("OpenAI-Beta") + } + req, err := s.buildUpstreamRequestOpenAIPassthrough(ctx, c, account, body, token) + if err != nil { + return nil, err + } + + targetURL, err := s.openAIAlphaSearchURL(account) + if err != nil { + return nil, err + } + parsedURL, err := url.Parse(targetURL) + if err != nil { + return nil, fmt.Errorf("parse alpha search URL: %w", err) + } + if c != nil && c.Request != nil && c.Request.URL != nil { + query := parsedURL.Query() + for key, values := range c.Request.URL.Query() { + for _, value := range values { + query.Add(key, value) + } + } + parsedURL.RawQuery = query.Encode() + } + req.URL = parsedURL + req.Header.Set("Accept", "application/json") + if clientBeta == "" { + req.Header.Del("OpenAI-Beta") + } + if version := strings.TrimSpace(c.GetHeader("Version")); version != "" { + req.Header.Set("Version", version) + } else if account.Type == AccountTypeOAuth { + req.Header.Set("Version", codexCLIVersion) + } + return req, nil +} + +func (s *OpenAIGatewayService) openAIAlphaSearchURL(account *Account) (string, error) { + if account == nil { + return "", fmt.Errorf("account is required") + } + switch account.Type { + case AccountTypeOAuth: + return chatgptCodexAlphaSearchURL, nil + case AccountTypeAPIKey: + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + return openAIPlatformAlphaSearchURL, nil + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return "", err + } + return buildOpenAIEndpointURL(validatedURL, "/v1/alpha/search"), nil + default: + return "", fmt.Errorf("unsupported OpenAI account type: %s", account.Type) + } +} diff --git a/backend/internal/service/openai_alpha_search_billing_test.go b/backend/internal/service/openai_alpha_search_billing_test.go new file mode 100644 index 0000000000..7151725763 --- /dev/null +++ b/backend/internal/service/openai_alpha_search_billing_test.go @@ -0,0 +1,101 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func TestCalculateWebSearchCostDefaultAndOverride(t *testing.T) { + t.Parallel() + s := &BillingService{} + + // 默认价:官方 $10/1000 次 = 0.01/次 + cost := s.CalculateWebSearchCost(1, nil, 1.0) + require.InDelta(t, 0.01, cost.TotalCost, 1e-12) + require.InDelta(t, 0.01, cost.ActualCost, 1e-12) + require.Equal(t, string(BillingModePerRequest), cost.BillingMode) + + // 分组覆盖价 + 倍率 + cost = s.CalculateWebSearchCost(1, float64Ptr(0.02), 2.5) + require.InDelta(t, 0.02, cost.TotalCost, 1e-12) + require.InDelta(t, 0.05, cost.ActualCost, 1e-12) + + // 0 = 免费(区别于 nil = 默认价) + cost = s.CalculateWebSearchCost(1, float64Ptr(0), 3.0) + require.Zero(t, cost.TotalCost) + require.Zero(t, cost.ActualCost) + + // 负数倍率按 0 处理,避免按 1x 误扣 + cost = s.CalculateWebSearchCost(1, nil, -1) + require.InDelta(t, 0.01, cost.TotalCost, 1e-12) + require.Zero(t, cost.ActualCost) + + // 次数 <= 0 不产生费用 + cost = s.CalculateWebSearchCost(0, float64Ptr(0.02), 1.0) + require.Zero(t, cost.TotalCost) + require.Empty(t, cost.BillingMode) +} + +func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { + t.Parallel() + svc := &OpenAIGatewayService{billingService: &BillingService{}} + groupID := int64(11) + + // 分组未配置单价:默认 0.01。按次搜索使用不含高峰因子的基础倍率(第 4 个倍率参数 2.0), + // 即使 token 倍率(含高峰,3.0)更高也不采用。 + apiKey := &APIKey{ID: 1, GroupID: &groupID, Group: &Group{ID: groupID, Platform: PlatformOpenAI}} + result := &OpenAIForwardResult{Model: "gpt-5.6-sol", UpstreamModel: "gpt-5.6-sol", WebSearchCalls: 1} + cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", false) + require.NoError(t, err) + require.Equal(t, string(BillingModePerRequest), cost.BillingMode) + require.InDelta(t, 0.01, cost.TotalCost, 1e-12) + require.InDelta(t, 0.02, cost.ActualCost, 1e-12) + + // 分组配置单价 0.005 + apiKey.Group.WebSearchPricePerCall = float64Ptr(0.005) + cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", false) + require.NoError(t, err) + require.InDelta(t, 0.005, cost.TotalCost, 1e-12) + require.InDelta(t, 0.005, cost.ActualCost, 1e-12) + + // WebSearchCalls = 0 时不得走按次分支(无定价数据会返回 pricing 错误, + // 证明回落到了 token 路径而不是被按次分支吞掉)。 + result.WebSearchCalls = 0 + _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", false) + require.Error(t, err) +} + +func TestAPIKeyService_SnapshotRoundTrip_PreservesWebSearchPricePerCall(t *testing.T) { + svc := NewAPIKeyService(nil, nil, nil, nil, nil, nil, &config.Config{}) + groupID := int64(9) + apiKey := &APIKey{ + ID: 1, + UserID: 2, + GroupID: &groupID, + Key: "k-websearch", + Status: StatusActive, + User: &User{ID: 2, Status: StatusActive, Role: RoleUser}, + Group: &Group{ + ID: groupID, + Name: "openai", + Platform: PlatformOpenAI, + Status: StatusActive, + SubscriptionType: SubscriptionTypeStandard, + RateMultiplier: 1, + WebSearchPricePerCall: float64Ptr(0.008), + }, + } + + snapshot := svc.snapshotFromAPIKey(context.Background(), apiKey) + roundTrip := svc.snapshotToAPIKey(apiKey.Key, snapshot) + + require.NotNil(t, roundTrip) + require.NotNil(t, roundTrip.Group) + require.NotNil(t, roundTrip.Group.WebSearchPricePerCall) + require.InDelta(t, 0.008, *roundTrip.Group.WebSearchPricePerCall, 1e-12) +} diff --git a/backend/internal/service/openai_alpha_search_test.go b/backend/internal/service/openai_alpha_search_test.go new file mode 100644 index 0000000000..458dbe9b5b --- /dev/null +++ b/backend/internal/service/openai_alpha_search_test.go @@ -0,0 +1,145 @@ +package service + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestForwardAlphaSearchOAuthPreservesWire(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{ + "id":"search-session", + "model":"gpt-5.6-sol", + "reasoning":{"effort":"max","context":"all_turns"}, + "input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"latest news"}]}], + "commands":{"search_query":[{"q":"OpenAI news","recency":1}]}, + "settings":{"allowed_callers":["direct"],"external_web_access":true}, + "max_output_tokens":2000, + "future_field":{"keep":true} + }`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search?feature=standalone", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", codexCLIUserAgent) + c.Request.Header.Set("Originator", "codex_cli_rs") + c.Request.Header.Set("Version", "0.144.1") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"encrypted_output":"ciphertext","output":"search result"}`)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 42, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-account", + }, + } + + result, err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, result.WebSearchCalls) + require.Equal(t, "gpt-5.6-sol", result.Model) + require.Equal(t, http.StatusOK, recorder.Code) + require.JSONEq(t, `{"encrypted_output":"ciphertext","output":"search result"}`, recorder.Body.String()) + require.Equal(t, chatgptCodexAlphaSearchURL+"?feature=standalone", upstream.lastReq.URL.String()) + require.Equal(t, "chatgpt.com", upstream.lastReq.Host) + require.Equal(t, "Bearer oauth-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "chatgpt-account", upstream.lastReq.Header.Get("chatgpt-account-id")) + require.Equal(t, "application/json", upstream.lastReq.Header.Get("Accept")) + require.Equal(t, "0.144.1", upstream.lastReq.Header.Get("Version")) + require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) + require.JSONEq(t, string(body), string(upstream.lastBody)) +} + +func TestForwardAlphaSearchAPIKeyMapsModelAndPassesThroughError(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{"search_query":[{"q":"news"}]}}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/alpha/search", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := `{"error":{"type":"invalid_request_error","message":"bad search"}}` + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 7, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://compat.example/v4", + "model_mapping": map[string]any{ + "gpt-5.6-sol": "upstream-5.6", + }, + }, + } + + result, err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + require.NoError(t, err) + // 上游错误透传不是一次成功的搜索:不返回 result、不产生按次计费。 + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, recorder.Code) + require.JSONEq(t, upstreamBody, recorder.Body.String()) + require.Equal(t, "https://compat.example/v4/alpha/search", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "upstream-5.6", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "commands.search_query").IsArray()) +} + +func TestForwardAlphaSearchReturnsFailoverBeforeWriting(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{}}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", bytes.NewReader(body)) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 8, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + }, + } + + result, err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, openAIPlatformAlphaSearchURL, upstream.lastReq.URL.String()) + require.False(t, c.Writer.Written()) + require.Empty(t, recorder.Body.String()) +} diff --git a/backend/internal/service/openai_codex_function_call_id_test.go b/backend/internal/service/openai_codex_function_call_id_test.go index 2ac59e0520..edfcac1d71 100644 --- a/backend/internal/service/openai_codex_function_call_id_test.go +++ b/backend/internal/service/openai_codex_function_call_id_test.go @@ -114,14 +114,15 @@ func TestFilterCodexInput_OutputTypeKeepsItemID(t *testing.T) { require.Equal(t, "o1", out["id"], "output item id should be preserved") } -// TestFilterCodexInput_NonToolCallItemKeepsID ensures non-tool-call items -// (e.g. message) still keep their id when PreserveReferences is true. +// TestFilterCodexInput_NonToolCallItemKeepsID ensures items subject to neither +// the fc* (call-input) nor the msg* (message) prefix rule still keep their id +// when PreserveReferences is true. +// message is covered separately in openai_codex_message_item_id_test.go (#3981). func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) { input := []any{ map[string]any{ - "type": "message", - "id": "item_msg_001", - "role": "user", + "type": "web_search_call", + "id": "ws_001", }, } @@ -130,7 +131,7 @@ func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) { }) require.Len(t, filtered, 1) - msg, ok := filtered[0].(map[string]any) + item, ok := filtered[0].(map[string]any) require.True(t, ok) - require.Equal(t, "item_msg_001", msg["id"], "non-tool-call items keep their id in preserve mode") + require.Equal(t, "ws_001", item["id"], "unconstrained items keep their id in preserve mode") } diff --git a/backend/internal/service/openai_codex_identity.go b/backend/internal/service/openai_codex_identity.go index 68d4105d84..47bc19ce0c 100644 --- a/backend/internal/service/openai_codex_identity.go +++ b/backend/internal/service/openai_codex_identity.go @@ -11,12 +11,32 @@ import ( // 若请求携带 version 且低于该值,上游直接 404(issue #3901,2026-07 实测)。 const codexUpstreamMinVersion = "0.144.0" +// ensureCodexIdentityHeaders 补齐 OAuth(ChatGPT 内部接口)出站请求所需的 Codex 身份头。 +// 已有 User-Agent 与 version 保持不变,交给紧随其后的 enforceCodexIdentityHeaders +// 做官方身份配对与最低版本校正。 +func ensureCodexIdentityHeaders(h http.Header) { + if h == nil { + return + } + if strings.TrimSpace(h.Get("user-agent")) == "" { + h.Set("user-agent", codexCLIUserAgent) + } + if strings.TrimSpace(h.Get("originator")) == "" { + h.Set("originator", "codex_cli_rs") + } + if strings.TrimSpace(h.Get("version")) == "" { + h.Set("version", codexCLIVersion) + } + h.Set("OpenAI-Beta", "responses=experimental") +} + // enforceCodexIdentityHeaders 收口 OAuth(ChatGPT 内部接口)出站请求的客户端身份头。 // 上游要求 originator 与 User-Agent 首段配套且为官方客户端标识,version 头(若携带) // 不低于 0.144.0,任一不满足即 404(issue #3901)。以最终 User-Agent 为准推导配套 // originator;推导不出官方身份(第三方 UA / UA 缺失)时整体回退为默认 Codex CLI 身份。 // -// 仅对携带 originator 的请求生效——compat messages bridge 故意不带 originator,保持原样。 +// 仅对携带 originator 的请求生效;需要从缺失身份头恢复的调用方应先调用 +// ensureCodexIdentityHeaders。 // 必须在所有 User-Agent 改写(自定义 UA / ForceCodexCLI / 浏览器 UA 兜底)之后调用。 func enforceCodexIdentityHeaders(h http.Header) { if h == nil || h.Get("originator") == "" { diff --git a/backend/internal/service/openai_codex_identity_test.go b/backend/internal/service/openai_codex_identity_test.go index 7d2c6d8520..ecb6eb7f6f 100644 --- a/backend/internal/service/openai_codex_identity_test.go +++ b/backend/internal/service/openai_codex_identity_test.go @@ -7,6 +7,36 @@ import ( "github.com/stretchr/testify/require" ) +func TestEnsureCodexIdentityHeaders(t *testing.T) { + t.Run("补齐缺失身份头", func(t *testing.T) { + h := make(http.Header) + + ensureCodexIdentityHeaders(h) + enforceCodexIdentityHeaders(h) + + require.Equal(t, "codex_cli_rs", h.Get("originator")) + require.Equal(t, codexCLIUserAgent, h.Get("user-agent")) + require.Equal(t, codexCLIVersion, h.Get("version")) + require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta")) + }) + + t.Run("保留已有官方UA和合法version并重新配对", func(t *testing.T) { + const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)" + h := make(http.Header) + h.Set("user-agent", tuiUA) + h.Set("version", "9.9.9") + h.Set("OpenAI-Beta", "assistants=v2") + + ensureCodexIdentityHeaders(h) + enforceCodexIdentityHeaders(h) + + require.Equal(t, "codex-tui", h.Get("originator")) + require.Equal(t, tuiUA, h.Get("user-agent")) + require.Equal(t, "9.9.9", h.Get("version")) + require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta")) + }) +} + func TestEnforceCodexIdentityHeaders(t *testing.T) { const tuiUA = "codex-tui/0.140.2 (Mac OS X 14.0; arm64) iTerm (codex-tui; 0.140.2)" @@ -102,13 +132,14 @@ func TestEnforceCodexIdentityHeaders(t *testing.T) { } } -// compat messages bridge 故意不带 originator:收口必须保持 no-op,不得注入身份头。 +// enforce 本身仍只负责收口:缺少 originator 时必须保持 no-op,由需要恢复身份的 +// 调用方先显式调用 ensureCodexIdentityHeaders。 func TestEnforceCodexIdentityHeaders_NoOriginatorIsNoop(t *testing.T) { h := make(http.Header) - h.Set("user-agent", "luna/1.0.0") + h.Set("user-agent", "third-party-client/1.0.0") enforceCodexIdentityHeaders(h) require.Empty(t, h.Get("originator")) - require.Equal(t, "luna/1.0.0", h.Get("user-agent")) + require.Equal(t, "third-party-client/1.0.0", h.Get("user-agent")) } diff --git a/backend/internal/service/openai_codex_message_item_id_test.go b/backend/internal/service/openai_codex_message_item_id_test.go new file mode 100644 index 0000000000..54c2a9ef22 --- /dev/null +++ b/backend/internal/service/openai_codex_message_item_id_test.go @@ -0,0 +1,160 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +// TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences +// verifies that message items with a non-msg id (e.g. item_*) have their id +// stripped even when PreserveReferences is true. OpenAI upstream requires +// message ids to begin with "msg" and rejects item_* with 400: +// "Expected an ID that begins with 'msg'." (#3981) +func TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "item_3bc5a3fa8ccde25f1c0000d4", + "role": "user", + "content": []any{ + map[string]any{"type": "input_text", "text": "hello"}, + }, + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "message", msg["type"]) + _, hasID := msg["id"] + require.False(t, hasID, "item_* id should be stripped from message") + require.Equal(t, "user", msg["role"], "role must be preserved") + require.NotNil(t, msg["content"], "content must be preserved") +} + +// TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences +// verifies that message items with a valid msg* id are kept when +// PreserveReferences is true, so context references are not lost. +func TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "msg_validID123", + "role": "assistant", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "msg_validID123", msg["id"], "valid msg* id must be preserved") +} + +// TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences ensures the +// non-continuation path still drops every message id regardless of prefix. +func TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences(t *testing.T) { + for _, id := range []string{"item_abc", "msg_validID123"} { + input := []any{ + map[string]any{ + "type": "message", + "id": id, + "role": "user", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: false, + }) + + require.Len(t, filtered, 1) + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + _, hasID := msg["id"] + require.False(t, hasID, "id %q should be stripped when not preserving references", id) + } +} + +// TestFilterCodexInput_MessageIDStripDoesNotMutateInput ensures the original +// input map is not modified in place when the id is stripped. +func TestFilterCodexInput_MessageIDStripDoesNotMutateInput(t *testing.T) { + original := map[string]any{ + "type": "message", + "id": "item_abc", + "role": "user", + } + + filtered := filterCodexInputWithOptions([]any{original}, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + require.Equal(t, "item_abc", original["id"], "original input must not be mutated") +} + +// TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior guards against a +// regression of #3785: message and function_call id rules are independent. +func TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "item_msg_001", + "role": "user", + }, + map[string]any{ + "type": "function_call", + "id": "fc_validID123", + "call_id": "fc_validID123", + "name": "bash", + }, + map[string]any{ + "type": "function_call", + "id": "item_A9v0SNfS3VaLrfX0j3y4xhyK", + "call_id": "fc_abc123", + "name": "bash", + }, + map[string]any{ + "type": "function_call_output", + "id": "o1", + "call_id": "fc_abc123", + "output": "done", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 4) + + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + _, hasID := msg["id"] + require.False(t, hasID, "message item_* id should be stripped") + + fcValid, ok := filtered[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "fc_validID123", fcValid["id"], "valid fc* id must be preserved") + + fcBad, ok := filtered[2].(map[string]any) + require.True(t, ok) + _, hasID = fcBad["id"] + require.False(t, hasID, "function_call item_* id should still be stripped") + require.Equal(t, "fc_abc123", fcBad["call_id"], "call_id pairing must survive") + + out, ok := filtered[3].(map[string]any) + require.True(t, ok) + require.Equal(t, "o1", out["id"], "output item id should be preserved") + require.Equal(t, "fc_abc123", out["call_id"], "call_id pairing must survive") +} diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index 0eb2e1a071..1d9356441c 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -2,21 +2,36 @@ package service import ( "context" + "crypto/sha256" + "errors" + "fmt" "io" + "net" "net/http" "net/url" + "sort" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "golang.org/x/net/http2" + "golang.org/x/sync/singleflight" ) // chatgptCodexModelsURL is the ChatGPT Codex models manifest endpoint. // Package-level variable so tests can point it at a stub server. var chatgptCodexModelsURL = "https://chatgpt.com/backend-api/codex/models" -const codexModelsManifestBodyLimit int64 = 8 << 20 +const ( + codexModelsManifestBodyLimit int64 = 8 << 20 + codexModelsManifestCacheBodyLimit = 1 << 20 + codexModelsManifestCacheMaxEntries = 64 + codexModelsManifestCacheTTL = 30 * time.Second + codexModelsManifestCacheStaleTTL = 5 * time.Minute + codexModelsManifestRequestTimeout = 15 * time.Second +) // CodexModelsManifest carries the raw upstream manifest payload plus caching // metadata so handlers can pass both through to the client untouched. @@ -26,8 +41,180 @@ type CodexModelsManifest struct { NotModified bool } -// FetchCodexModelsManifest fetches the live Codex models manifest from the -// ChatGPT backend using the account's OAuth credentials. +type codexModelsManifestUpstreamError struct { + err error + retryable bool +} + +func (e *codexModelsManifestUpstreamError) Error() string { return e.err.Error() } + +func (e *codexModelsManifestUpstreamError) Unwrap() error { return e.err } + +// IsRetryableCodexModelsManifestError reports whether another selected account +// may succeed without changing the request. Configuration and upstream 4xx +// responses, except 429, are intentionally not retried. +func IsRetryableCodexModelsManifestError(err error) bool { + var upstreamErr *codexModelsManifestUpstreamError + return errors.As(err, &upstreamErr) && upstreamErr.retryable +} + +func isRetryableCodexModelsManifestTransportError(err error) bool { + if err == nil || errors.Is(err, context.Canceled) { + return false + } + if errors.Is(err, context.DeadlineExceeded) || + errors.Is(err, io.EOF) || + errors.Is(err, io.ErrUnexpectedEOF) || + errors.Is(err, net.ErrClosed) { + return true + } + + var opErr *net.OpError + if errors.As(err, &opErr) { + return true + } + var dnsErr *net.DNSError + if errors.As(err, &dnsErr) { + return true + } + var goAwayErr http2.GoAwayError + if errors.As(err, &goAwayErr) { + return true + } + var streamErr http2.StreamError + if errors.As(err, &streamErr) { + return true + } + var connectionErr http2.ConnectionError + if errors.As(err, &connectionErr) { + return true + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + + // net/http uses unexported HTTP/2 error types, so typed matching is not + // possible for errors produced by the standard library transport. + message := strings.ToLower(err.Error()) + if strings.Contains(message, "http2:") && + (strings.Contains(message, "goaway") || + strings.Contains(message, "refused_stream") || + strings.Contains(message, "frame too large")) { + return true + } + if strings.Contains(message, "stream error: stream id ") { + return true + } + for _, code := range []http2.ErrCode{ + http2.ErrCodeNo, + http2.ErrCodeProtocol, + http2.ErrCodeInternal, + http2.ErrCodeFlowControl, + http2.ErrCodeSettingsTimeout, + http2.ErrCodeStreamClosed, + http2.ErrCodeFrameSize, + http2.ErrCodeRefusedStream, + http2.ErrCodeCancel, + http2.ErrCodeCompression, + http2.ErrCodeConnect, + http2.ErrCodeEnhanceYourCalm, + http2.ErrCodeInadequateSecurity, + http2.ErrCodeHTTP11Required, + } { + if strings.Contains(message, "connection error: "+strings.ToLower(code.String())) { + return true + } + } + return false +} + +type codexModelsManifestRequest struct { + url string + headers http.Header + proxyURL string + accountID int64 + credentialAccountID int64 + accountConcurrency int + useAPIKeyUpstream bool +} + +type codexModelsManifestCacheEntry struct { + manifest *CodexModelsManifest + order uint64 + expiresAt time.Time + staleUntil time.Time +} + +type codexModelsManifestCacheState uint8 + +const ( + codexModelsManifestCacheMiss codexModelsManifestCacheState = iota + codexModelsManifestCacheFresh + codexModelsManifestCacheStale +) + +type codexModelsManifestCache struct { + mu sync.Mutex + entries map[string]codexModelsManifestCacheEntry + nextOrder uint64 + refresh singleflight.Group +} + +func (c *codexModelsManifestCache) get(key string, now time.Time) (*CodexModelsManifest, codexModelsManifestCacheState) { + c.mu.Lock() + defer c.mu.Unlock() + entry, ok := c.entries[key] + if !ok { + return nil, codexModelsManifestCacheMiss + } + if !now.Before(entry.staleUntil) { + delete(c.entries, key) + return nil, codexModelsManifestCacheMiss + } + if now.Before(entry.expiresAt) { + return entry.manifest, codexModelsManifestCacheFresh + } + return entry.manifest, codexModelsManifestCacheStale +} + +func (c *codexModelsManifestCache) set(key string, manifest *CodexModelsManifest, now time.Time) { + if manifest == nil || len(manifest.Body) > codexModelsManifestCacheBodyLimit { + return + } + c.mu.Lock() + defer c.mu.Unlock() + if c.entries == nil { + c.entries = make(map[string]codexModelsManifestCacheEntry) + } + if _, exists := c.entries[key]; !exists && len(c.entries) >= codexModelsManifestCacheMaxEntries { + oldestKey := "" + var oldestOrder uint64 + for candidateKey, entry := range c.entries { + if !now.Before(entry.staleUntil) { + delete(c.entries, candidateKey) + continue + } + if oldestKey == "" || entry.order < oldestOrder { + oldestKey = candidateKey + oldestOrder = entry.order + } + } + if len(c.entries) >= codexModelsManifestCacheMaxEntries && oldestKey != "" { + delete(c.entries, oldestKey) + } + } + c.nextOrder++ + c.entries[key] = codexModelsManifestCacheEntry{ + manifest: manifest, + order: c.nextOrder, + expiresAt: now.Add(codexModelsManifestCacheTTL), + staleUntil: now.Add(codexModelsManifestCacheStaleTTL), + } +} + +// FetchCodexModelsManifest fetches the live Codex models manifest from either +// the ChatGPT backend for OAuth accounts or a custom upstream for API key accounts. // // The response body is passed through verbatim: the manifest schema evolves // with Codex client releases, and interpreting it here would force the gateway @@ -41,57 +228,180 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if err != nil { return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_CREDENTIALS_FAILED", "resolve credential account: %v", err) } - accessToken := credAccount.GetOpenAIAccessToken() - if accessToken == "" && !credAccount.IsOpenAIAgentIdentity() { - return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") - } clientVersion = strings.TrimSpace(clientVersion) if clientVersion == "" { clientVersion = openAICodexProbeVersion } - requestURL := chatgptCodexModelsURL + "?client_version=" + url.QueryEscape(clientVersion) - reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second) - defer cancel() - req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL, nil) - if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) - } - authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, credAccount, accessToken) - if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_AUTH_FAILED", "build Codex models authentication: %v", err) - } - for key, values := range authHeaders { - for _, value := range values { - req.Header.Add(key, value) + requestEndpoint := chatgptCodexModelsURL + authToken := "" + useAPIKeyUpstream := false + appendModelsPath := false + switch { + case credAccount.IsOpenAIOAuth(): + authToken = strings.TrimSpace(credAccount.GetOpenAIAccessToken()) + if authToken == "" && !credAccount.IsOpenAIAgentIdentity() { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") } + case credAccount.IsOpenAIApiKey(): + baseURL := strings.TrimSpace(credAccount.GetCredential("base_url")) + if baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) { + return nil, infraerrors.New( + http.StatusBadGateway, + "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED", + "Codex models manifest requires a custom API key upstream base URL", + ) + } + authToken = strings.TrimSpace(credAccount.GetOpenAIApiKey()) + if authToken == "" { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_MISSING", "account has no API key for the Codex models upstream") + } + normalizedBaseURL, validateErr := s.validateUpstreamBaseURL(baseURL) + if validateErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", validateErr) + } + requestEndpoint = normalizedBaseURL + useAPIKeyUpstream = true + appendModelsPath = true + default: + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_ACCOUNT_TYPE_UNSUPPORTED", "account type %q cannot fetch the Codex models manifest", credAccount.Type) } - req.Header.Set("Accept", "application/json") - req.Header.Set("Originator", "codex_cli_rs") - req.Header.Set("Version", clientVersion) - req.Header.Set("User-Agent", codexCLIUserAgent) - if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { - req.Header.Set("If-None-Match", ifNoneMatch) + + requestURL, err := buildCodexModelsManifestURL(requestEndpoint, appendModelsPath, clientVersion) + if err != nil { + if useAPIKeyUpstream { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", err) + } + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "parse codex models request URL: %v", err) } - setOpenAIChatGPTAccountHeaders(req.Header, credAccount) + + headers := make(http.Header) + if useAPIKeyUpstream { + headers.Set("Authorization", "Bearer "+authToken) + credAccount.ApplyHeaderOverrides(headers) + } else { + authHeaders, authErr := s.buildOpenAIAuthenticationHeaders(ctx, credAccount, authToken) + if authErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_AUTH_FAILED", "build Codex models authentication: %v", authErr) + } + for key, values := range authHeaders { + for _, value := range values { + headers.Add(key, value) + } + } + setOpenAIChatGPTAccountHeaders(headers, credAccount) + } + headers.Set("Accept", "application/json") + headers.Set("Originator", "codex_cli_rs") + headers.Set("Version", clientVersion) + headers.Set("User-Agent", codexCLIUserAgent) proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() } - client, err := httpclient.GetClient(httpclient.Options{ - ProxyURL: proxyURL, - Timeout: 15 * time.Second, - ResponseHeaderTimeout: 10 * time.Second, + + request := codexModelsManifestRequest{ + url: requestURL.String(), + headers: headers, + proxyURL: proxyURL, + accountID: account.ID, + credentialAccountID: credAccount.ID, + accountConcurrency: account.Concurrency, + useAPIKeyUpstream: useAPIKeyUpstream, + } + if useAPIKeyUpstream { + return s.fetchCachedAPIKeyCodexModelsManifest(ctx, request, ifNoneMatch) + } + return s.fetchCodexModelsManifestUpstream(ctx, request, ifNoneMatch) +} + +func (s *OpenAIGatewayService) fetchCachedAPIKeyCodexModelsManifest(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + cacheKey := buildCodexModelsManifestCacheKey(request) + manifest, state := s.codexModelsManifestCache.get(cacheKey, time.Now()) + if state == codexModelsManifestCacheFresh { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + resultCh := s.refreshCachedAPIKeyCodexModelsManifest(cacheKey, request) + if state == codexModelsManifestCacheStale { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-resultCh: + if result.Err != nil { + return nil, result.Err + } + manifest, ok := result.Val.(*CodexModelsManifest) + if !ok || manifest == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "invalid shared Codex models manifest result") + } + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } +} + +func (s *OpenAIGatewayService) refreshCachedAPIKeyCodexModelsManifest(cacheKey string, request codexModelsManifestRequest) <-chan singleflight.Result { + return s.codexModelsManifestCache.refresh.DoChan(cacheKey, func() (any, error) { + cached, _ := s.codexModelsManifestCache.get(cacheKey, time.Now()) + ifNoneMatch := "" + if cached != nil { + ifNoneMatch = cached.ETag + } + manifest, err := s.fetchCodexModelsManifestUpstream(context.Background(), request, ifNoneMatch) + if err != nil { + return nil, err + } + if manifest.NotModified && cached != nil { + s.codexModelsManifestCache.set(cacheKey, cached, time.Now()) + return cached, nil + } + if !manifest.NotModified { + s.codexModelsManifestCache.set(cacheKey, manifest, time.Now()) + } + return manifest, nil }) +} + +func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + reqCtx, cancel := context.WithTimeout(ctx, codexModelsManifestRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, request.url, nil) if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", err) + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) + } + req.Header = request.headers.Clone() + if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { + req.Header.Set("If-None-Match", ifNoneMatch) } - resp, err := client.Do(req) + var resp *http.Response + if request.useAPIKeyUpstream { + if s.httpUpstream == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_UPSTREAM_NOT_CONFIGURED", "Codex models upstream HTTP client is not configured") + } + req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) + resp, err = s.httpUpstream.Do(req, request.proxyURL, request.accountID, request.accountConcurrency) + } else { + client, clientErr := httpclient.GetClient(httpclient.Options{ + ProxyURL: request.proxyURL, + Timeout: codexModelsManifestRequestTimeout, + ResponseHeaderTimeout: 10 * time.Second, + }) + if clientErr != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", clientErr) + } + resp, err = client.Do(req) + } if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } defer func() { _ = resp.Body.Close() }() @@ -104,12 +414,100 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if message == "" { message = resp.Status } - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message), + retryable: resp.StatusCode == http.StatusTooManyRequests || + (resp.StatusCode >= http.StatusInternalServerError && resp.StatusCode < 600), + } } body, err := io.ReadAll(io.LimitReader(resp.Body, codexModelsManifestBodyLimit)) if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil } + +func buildCodexModelsManifestCacheKey(request codexModelsManifestRequest) string { + hasher := sha256.New() + _, _ = fmt.Fprintf(hasher, "%d\n%d\n%s\n%s\n", request.accountID, request.credentialAccountID, request.proxyURL, request.url) + headerNames := make([]string, 0, len(request.headers)) + for name := range request.headers { + headerNames = append(headerNames, name) + } + sort.Strings(headerNames) + for _, name := range headerNames { + _, _ = fmt.Fprintf(hasher, "%s\n", strings.ToLower(name)) + for _, value := range request.headers[name] { + _, _ = fmt.Fprintf(hasher, "%s\n", value) + } + } + return fmt.Sprintf("%x", hasher.Sum(nil)) +} + +func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch string) *CodexModelsManifest { + if manifest == nil { + return nil + } + if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { + return &CodexModelsManifest{ETag: manifest.ETag, NotModified: true} + } + return manifest +} + +func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool { + etag = strings.TrimSpace(etag) + if etag == "" { + return false + } + normalize := func(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && strings.EqualFold(value[:2], "W/") { + value = strings.TrimSpace(value[2:]) + } + return value + } + want := normalize(etag) + for _, candidate := range strings.Split(ifNoneMatch, ",") { + candidate = strings.TrimSpace(candidate) + if candidate == "*" || normalize(candidate) == want { + return true + } + } + return false +} + +func isOfficialOpenAIModelsBaseURL(raw string) bool { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil { + return false + } + hostname := strings.TrimSuffix(parsed.Hostname(), ".") + return strings.EqualFold(hostname, "api.openai.com") +} + +func buildCodexModelsManifestURL(endpoint string, appendModelsPath bool, clientVersion string) (*url.URL, error) { + requestURL, err := url.Parse(endpoint) + if err != nil { + return nil, err + } + if requestURL.Fragment != "" { + return nil, fmt.Errorf("URL fragments are not supported") + } + + query := requestURL.Query() + requestURL.RawQuery = "" + requestURL.ForceQuery = false + if appendModelsPath { + requestURL, err = url.Parse(buildOpenAIModelsURL(requestURL.String())) + if err != nil { + return nil, err + } + } + query.Set("client_version", clientVersion) + requestURL.RawQuery = query.Encode() + return requestURL, nil +} diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index c9eae35629..3f4ec96a2f 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -2,11 +2,146 @@ package service import ( "context" + "errors" + "io" + "net" "net/http" "net/http/httptest" + "net/url" + "strings" + "sync" + "sync/atomic" "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "golang.org/x/net/http2" ) +type codexModelsHTTPUpstreamStub struct { + do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) +} + +type codexModelsBlockingBody struct { + ctx context.Context + readStarted chan struct{} + startedOnce *sync.Once + release <-chan struct{} + body *strings.Reader +} + +func (b *codexModelsBlockingBody) Read(p []byte) (int, error) { + b.startedOnce.Do(func() { close(b.readStarted) }) + select { + case <-b.release: + return b.body.Read(p) + case <-b.ctx.Done(): + return 0, b.ctx.Err() + } +} + +func (b *codexModelsBlockingBody) Close() error { return nil } + +func (s *codexModelsHTTPUpstreamStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + return s.do(req, proxyURL, accountID, accountConcurrency) +} + +func (s *codexModelsHTTPUpstreamStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return s.Do(req, proxyURL, accountID, accountConcurrency) +} + +func TestIsRetryableCodexModelsManifestTransportError(t *testing.T) { + tests := []struct { + name string + err error + retryable bool + }{ + {name: "nil", err: nil}, + {name: "configuration error", err: errors.New("invalid proxy URL")}, + {name: "upstream configuration error", err: errors.New("upstream error: invalid proxy")}, + {name: "proxy connection configuration error", err: errors.New("proxy connection error: invalid configuration")}, + {name: "canceled request", err: context.Canceled}, + { + name: "redirect policy error", + err: &url.Error{ + Op: "Get", + URL: "https://upstream.example/v1/models", + Err: errors.New("stopped after 10 redirects"), + }, + }, + {name: "deadline exceeded", err: context.DeadlineExceeded, retryable: true}, + {name: "unexpected EOF", err: io.ErrUnexpectedEOF, retryable: true}, + {name: "closed connection", err: net.ErrClosed, retryable: true}, + { + name: "network operation", + err: &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + }, + retryable: true, + }, + { + name: "DNS error", + err: &net.DNSError{Err: "temporary failure", Name: "upstream.example"}, + retryable: true, + }, + { + name: "typed HTTP2 GOAWAY", + err: http2.GoAwayError{ErrCode: http2.ErrCodeNo}, + retryable: true, + }, + { + name: "stdlib HTTP2 GOAWAY", + err: errors.New("http2: server sent GOAWAY and closed the connection; LastStreamID=1, ErrCode=NO_ERROR"), + retryable: true, + }, + { + name: "stdlib HTTP2 refused stream", + err: errors.New("stream error: stream ID 3; REFUSED_STREAM"), + retryable: true, + }, + { + name: "stdlib HTTP2 connection error", + err: errors.New(`Get "https://upstream.example/v1/models": connection error: PROTOCOL_ERROR`), + retryable: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isRetryableCodexModelsManifestTransportError(tt.err); got != tt.retryable { + t.Fatalf("retryable = %v, want %v", got, tt.retryable) + } + }) + } +} + +func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{ + Enabled: false, + }}}, + httpUpstream: upstream, + } +} + +func newCodexModelsAPIKeyTestAccount(baseURL string) *Account { + credentials := map[string]any{"api_key": "sk-upstream"} + if baseURL != "" { + credentials["base_url"] = baseURL + } + return &Account{ + ID: 2, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: credentials, + Concurrency: 3, + } +} + func newCodexModelsTestAccount() *Account { return &Account{ ID: 1, @@ -64,6 +199,49 @@ func TestFetchCodexModelsManifestPassthrough(t *testing.T) { } } +func TestFetchCodexModelsManifestAgentIdentityUsesAssertionWithoutOAuthToken(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 3, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": key.taskID, + "chatgpt_account_id": "acc-agent", + }, + } + + var gotAuth, gotAccountID string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + gotAccountID = r.Header.Get("chatgpt-account-id") + _, _ = w.Write([]byte(`{"models":[]}`)) + })) + defer server.Close() + + original := chatgptCodexModelsURL + chatgptCodexModelsURL = server.URL + defer func() { chatgptCodexModelsURL = original }() + + s := &OpenAIGatewayService{} + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.137.0", "") + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if string(manifest.Body) != `{"models":[]}` { + t.Fatalf("unexpected manifest body: %q", manifest.Body) + } + if !strings.HasPrefix(gotAuth, "AgentAssertion ") { + t.Fatalf("authorization scheme: got %q", strings.SplitN(gotAuth, " ", 2)[0]) + } + if gotAccountID != "acc-agent" { + t.Fatalf("chatgpt-account-id header: got %q", gotAccountID) + } +} + func TestFetchCodexModelsManifestDefaultClientVersion(t *testing.T) { var gotClientVersion string server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -136,3 +314,679 @@ func TestFetchCodexModelsManifestMissingToken(t *testing.T) { t.Fatal("expected error for missing access token, got nil") } } + +func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { + manifestBody := `{"models":[{"slug":"gpt-5.6"}]}` + var gotRequest *http.Request + var gotProxyURL string + var gotAccountID int64 + var gotConcurrency int + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + gotRequest = req + gotProxyURL = proxyURL + gotAccountID = accountID + gotConcurrency = accountConcurrency + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(manifestBody)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + + if gotRequest == nil { + t.Fatal("expected request to custom API key upstream") + } + if gotRequest.Method != http.MethodGet { + t.Errorf("method: got %q", gotRequest.Method) + } + if gotRequest.URL.String() != "https://upstream.example/v1/models?client_version=0.144.0" { + t.Errorf("request URL: got %q", gotRequest.URL.String()) + } + if gotRequest.Header.Get("Authorization") != "Bearer sk-upstream" { + t.Errorf("authorization header: got %q", gotRequest.Header.Get("Authorization")) + } + if gotRequest.Header.Get("Originator") != "codex_cli_rs" { + t.Errorf("originator header: got %q", gotRequest.Header.Get("Originator")) + } + if gotRequest.Header.Get("Version") != "0.144.0" { + t.Errorf("version header: got %q", gotRequest.Header.Get("Version")) + } + if gotRequest.Header.Get("User-Agent") != codexCLIUserAgent { + t.Errorf("user-agent header: got %q", gotRequest.Header.Get("User-Agent")) + } + if gotRequest.Header.Get("chatgpt-account-id") != "" { + t.Errorf("chatgpt-account-id must not be sent to API key upstream: got %q", gotRequest.Header.Get("chatgpt-account-id")) + } + if gotProxyURL != "" || gotAccountID != 2 || gotConcurrency != 3 { + t.Errorf("upstream routing metadata: proxy=%q account_id=%d concurrency=%d", gotProxyURL, gotAccountID, gotConcurrency) + } + if string(manifest.Body) != manifestBody { + t.Errorf("body not passed through verbatim: got %q", manifest.Body) + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } +} + +func TestFetchCodexModelsManifestAPIKeySharedRefreshSurvivesCallerCancellation(t *testing.T) { + const manifestBody = `{"models":[{"slug":"gpt-5.6"}]}` + var calls atomic.Int32 + var readStartedOnce sync.Once + readStarted := make(chan struct{}) + deadlineRemaining := make(chan time.Duration, 1) + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + deadline, ok := req.Context().Deadline() + if !ok { + deadlineRemaining <- 0 + } else { + deadlineRemaining <- time.Until(deadline) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"shared"`}}, + Body: &codexModelsBlockingBody{ + ctx: req.Context(), + readStarted: readStarted, + startedOnce: &readStartedOnce, + release: release, + body: strings.NewReader(manifestBody), + }, + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + firstCtx, cancelFirst := context.WithCancel(context.Background()) + firstErr := make(chan error, 1) + go func() { + _, err := s.FetchCodexModelsManifest(firstCtx, account, "0.144.0", "") + firstErr <- err + }() + + select { + case <-readStarted: + case <-time.After(time.Second): + t.Fatal("upstream body read did not start") + } + remaining := <-deadlineRemaining + if remaining < 14*time.Second || remaining > codexModelsManifestRequestTimeout { + t.Errorf("detached refresh deadline: got %s, want approximately %s", remaining, codexModelsManifestRequestTimeout) + } + cancelFirst() + select { + case err := <-firstErr: + if !errors.Is(err, context.Canceled) { + t.Fatalf("first caller error: got %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("canceled caller did not return promptly") + } + + secondResult := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + secondResult <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls before shared refresh completed: got %d, want 1", got) + } + close(release) + select { + case result := <-secondResult: + if result.err != nil { + t.Fatalf("second caller returned error: %v", result.err) + } + if string(result.manifest.Body) != manifestBody { + t.Errorf("second caller body: got %q", result.manifest.Body) + } + case <-time.After(time.Second): + t.Fatal("second caller did not receive shared refresh result") + } + if got := calls.Load(); got != 1 { + t.Errorf("total upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyConcurrentRequestsShareRefresh(t *testing.T) { + const callers = 8 + var calls atomic.Int32 + started := make(chan struct{}) + var startedOnce sync.Once + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + startedOnce.Do(func() { close(started) }) + <-release + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + begin := make(chan struct{}) + errs := make(chan error, callers) + for i := 0; i < callers; i++ { + go func() { + <-begin + _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + errs <- err + }() + } + close(begin) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("upstream request did not start") + } + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("concurrent upstream calls: got %d, want 1", got) + } + close(release) + for i := 0; i < callers; i++ { + if err := <-errs; err != nil { + t.Errorf("caller %d returned error: %v", i, err) + } + } +} + +func TestFetchCodexModelsManifestAPIKeyFreshCacheHandlesETagLocally(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cache refresh must not inherit a caller's If-None-Match: got %q", got) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"cached"`}}, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", `W/"cached"`) + if err != nil { + t.Fatalf("cached fetch returned error: %v", err) + } + if !manifest.NotModified { + t.Fatal("matching cached ETag must return NotModified") + } + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheKeyIsolatesRequestIdentity(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + + base := newCodexModelsAPIKeyTestAccount("https://upstream.example") + fetch := func(account *Account, version string) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, version, ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + fetch(base, "0.144.0") + fetch(base, "0.144.0") + + differentAccount := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentAccount.ID = 3 + fetch(differentAccount, "0.144.0") + + differentToken := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentToken.Credentials["api_key"] = "sk-other" + fetch(differentToken, "0.144.0") + + differentUpstream := newCodexModelsAPIKeyTestAccount("https://other-upstream.example") + fetch(differentUpstream, "0.144.0") + fetch(base, "0.145.0") + + differentHeaders := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentHeaders.Credentials[credKeyHeaderOverrideEnabled] = true + differentHeaders.Credentials[credKeyHeaderOverrides] = map[string]any{"x-tenant": "other"} + fetch(differentHeaders, "0.144.0") + + proxyID := int64(9) + differentProxy := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentProxy.ProxyID = &proxyID + differentProxy.Proxy = &Proxy{Protocol: "http", Host: "127.0.0.1", Port: 8080} + fetch(differentProxy, "0.144.0") + fetch(differentProxy, "0.144.0") + + if got := calls.Load(); got != 7 { + t.Errorf("isolated upstream calls: got %d, want 7", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + body := `{"models":[]}` + if strings.Contains(req.URL.Host, "large") { + body = strings.Repeat("x", (1<<20)+1) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + fetch := func(account *Account) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + + small := newCodexModelsAPIKeyTestAccount("https://small.example") + fetch(small) + fetch(small) + large := newCodexModelsAPIKeyTestAccount("https://large.example") + large.ID = 3 + fetch(large) + fetch(large) + if got := calls.Load(); got != 3 { + t.Fatalf("body-size bounded cache calls: got %d, want 3", got) + } + + for i := int64(10); i < 75; i++ { + account := newCodexModelsAPIKeyTestAccount("https://bounded.example") + account.ID = i + fetch(account) + } + last := newCodexModelsAPIKeyTestAccount("https://bounded.example") + last.ID = 74 + fetch(last) + if got := calls.Load(); got != 68 { + t.Fatalf("most recent cache entry was not retained: calls=%d, want 68", got) + } + first := newCodexModelsAPIKeyTestAccount("https://bounded.example") + first.ID = 10 + fetch(first) + if got := calls.Load(); got != 69 { + t.Errorf("oldest cache entry was not evicted: calls=%d, want 69", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyServesStaleWhileRefreshing(t *testing.T) { + var calls atomic.Int32 + refreshStarted := make(chan struct{}) + releaseRefresh := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + body := `{"models":[{"slug":"old"}]}` + if call > 1 { + if call == 2 { + close(refreshStarted) + } + <-releaseRefresh + body = `{"models":[{"slug":"new"}]}` + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + resultCh := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + resultCh <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + select { + case <-refreshStarted: + case <-time.After(time.Second): + t.Fatal("background refresh did not start") + } + + var staleResult struct { + manifest *CodexModelsManifest + err error + } + select { + case staleResult = <-resultCh: + case <-time.After(100 * time.Millisecond): + t.Error("stale manifest was not returned while refresh was blocked") + close(releaseRefresh) + staleResult = <-resultCh + } + if staleResult.err != nil { + t.Fatalf("stale fetch returned error: %v", staleResult.err) + } + if got := string(staleResult.manifest.Body); got != `{"models":[{"slug":"old"}]}` { + t.Errorf("stale body: got %q", got) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls during stale refresh: got %d, want 2", got) + } + + select { + case <-releaseRefresh: + default: + close(releaseRefresh) + } + deadline := time.Now().Add(time.Second) + for { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err == nil && string(manifest.Body) == `{"models":[{"slug":"new"}]}` { + break + } + if time.Now().After(deadline) { + t.Fatalf("refreshed manifest was not cached: manifest=%v err=%v", manifest, err) + } + time.Sleep(10 * time.Millisecond) + } + if got := calls.Load(); got != 2 { + t.Errorf("stale refresh was not deduplicated: calls=%d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyRevalidatesStaleETag(t *testing.T) { + var calls atomic.Int32 + refreshDone := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + if call == 1 { + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"cached"}]}`)), + }, nil + } + if got := req.Header.Get("If-None-Match"); got != `W/"cached"` { + t.Errorf("background revalidation If-None-Match: got %q", got) + } + close(refreshDone) + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("stale fetch returned error: %v", err) + } + if got := string(manifest.Body); got != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("stale body: got %q", got) + } + select { + case <-refreshDone: + case <-time.After(time.Second): + t.Fatal("ETag revalidation did not complete") + } + + deadline := time.Now().Add(time.Second) + for { + s.codexModelsManifestCache.mu.Lock() + fresh := false + for _, entry := range s.codexModelsManifestCache.entries { + fresh = time.Now().Before(entry.expiresAt) + } + s.codexModelsManifestCache.mu.Unlock() + if fresh { + break + } + if time.Now().After(deadline) { + t.Fatal("304 revalidation did not renew the cached manifest") + } + time.Sleep(10 * time.Millisecond) + } + manifest, err = s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil || string(manifest.Body) != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("renewed cached manifest: body=%q err=%v", manifest.Body, err) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls: got %d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyColdCacheHandlesNotModifiedLocally(t *testing.T) { + var gotIfNoneMatch string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotIfNoneMatch = req.Header.Get("If-None-Match") + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + `W/"api-key-manifest"`, + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if !manifest.NotModified { + t.Error("expected NotModified to be true") + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } + if gotIfNoneMatch != "" { + t.Errorf("cold shared refresh must not inherit caller if-none-match: got %q", gotIfNoneMatch) + } +} + +func TestFetchCodexModelsManifestAPIKeyDoesNotCacheUnexpectedColdNotModified(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cold shared refresh If-None-Match: got %q", got) + } + header := make(http.Header) + header.Set("ETag", `W/"unexpected"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + for i := 0; i < 2; i++ { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("fetch %d returned error: %v", i, err) + } + if !manifest.NotModified { + t.Fatalf("fetch %d: expected upstream NotModified response", i) + } + } + if got := calls.Load(); got != 2 { + t.Errorf("unexpected cold 304 was cached: upstream calls=%d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyPreservesBaseURLQuery(t *testing.T) { + var gotURL string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotURL = req.URL.String() + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1?tenant=acme"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if gotURL != "https://upstream.example/v1/models?client_version=0.144.0&tenant=acme" { + t.Errorf("request URL: got %q", gotURL) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsBaseURLFragment(t *testing.T) { + called := false + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + called = true + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1#models"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected invalid upstream base URL error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + if called { + t.Fatal("fragment-bearing base URL must be rejected before the upstream request") + } +} + +func TestFetchCodexModelsManifestAPIKeyUpstreamError(t *testing.T) { + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusTooManyRequests, + Status: "429 Too Many Requests", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":"rate limited"}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected error for upstream 429, got nil") + } + if infraerrors.Code(err) != http.StatusBadGateway { + t.Errorf("error status: got %d, want %d", infraerrors.Code(err), http.StatusBadGateway) + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_UPSTREAM_FAILED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T) { + tests := []struct { + name string + baseURL string + }{ + {name: "missing base URL"}, + {name: "official host", baseURL: "https://api.openai.com"}, + {name: "official versioned URL", baseURL: "https://API.OPENAI.COM:443/v1/"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := newCodexModelsAPIKeyTestService(&codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + t.Fatal("official OpenAI API key must not be used as a Codex manifest upstream") + return nil, nil + }}) + + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount(tt.baseURL), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected unsupported API key upstream error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + }) + } +} diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 99355628f2..3bde4abcfb 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -838,6 +838,9 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool { if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) { return false } + if hasOpenAIImageGenerationTool(reqBody) { + return false + } tool := map[string]any{ "type": "image_generation", @@ -855,16 +858,6 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool { reqBody["tools"] = []any{tool} return true } - for _, rawTool := range tools { - toolMap, ok := rawTool.(map[string]any) - if !ok { - continue - } - if strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { - return false - } - } - reqBody["tools"] = append(tools, tool) return true } @@ -1405,6 +1398,15 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an ensureCopy() delete(newItem, "id") } + } else if typ == "message" { + // 同理,message 类 item 的 id 必须以 "msg" 开头(上游校验 + // "Expected an ID that begins with 'msg'")。item_* 形式的 id + // 来自客户端回放,需要删除。 + // 注意:不改写成 msg_*,改写出的 id 未必对应真实的上游对象。 + if id, ok := m["id"].(string); ok && id != "" && !strings.HasPrefix(id, "msg") { + ensureCopy() + delete(newItem, "id") + } } filtered = append(filtered, newItem) diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index b226655eeb..456740136b 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -617,6 +617,65 @@ func TestEnsureOpenAIResponsesImageGenerationTool_PreservesExistingImageTool(t * require.Equal(t, "webp", tool["output_format"]) } +func TestEnsureOpenAIResponsesImageGenerationTool_PreservesImageGenNamespace(t *testing.T) { + tests := []struct { + name string + reqBody map[string]any + }{ + { + name: "top-level tools", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + }, + { + name: "responses lite additional_tools", + reqBody: map[string]any{ + "model": "gpt-5.5", + "input": []any{ + map[string]any{ + "type": "additional_tools", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.True(t, hasOpenAIImageGenerationTool(tt.reqBody)) + + modified := ensureOpenAIResponsesImageGenerationTool(tt.reqBody) + + require.False(t, modified) + tools, _ := tt.reqBody["tools"].([]any) + for _, rawTool := range tools { + tool, ok := rawTool.(map[string]any) + require.True(t, ok) + require.NotEqual(t, "image_generation", firstNonEmptyString(tool["type"])) + } + }) + } +} + func TestApplyCodexImageGenerationBridgeInstructions_AppendsBridgeOnce(t *testing.T) { reqBody := map[string]any{ "model": "gpt-5.4", diff --git a/backend/internal/service/openai_compact_body_signal.go b/backend/internal/service/openai_compact_body_signal.go index fce62046c1..ce561b0c5a 100644 --- a/backend/internal/service/openai_compact_body_signal.go +++ b/backend/internal/service/openai_compact_body_signal.go @@ -2,18 +2,10 @@ package service import "github.com/tidwall/gjson" -// HasCompactionTriggerInInput detects the Codex remote compact v2 body signal: -// an input item with type "compaction_trigger". When the client sends this -// inside a normal POST /v1/responses (instead of POST /v1/responses/compact), -// the request must still be treated as a compact request — otherwise the -// upstream path, model mapping, and body normalization are all wrong, causing -// Codex to receive a non-compact response and fail with: -// -// "remote compaction v2 expected exactly one compaction output item, got 0" -// -// The gateway handler promotes such requests by rewriting the URL path to the -// compact form before stream parsing, compact body normalization, and -// compact-capable account scheduling, so both inbound forms share one code path. +// HasCompactionTriggerInInput detects an input item with +// type="compaction_trigger". The handler combines this body signal with the +// request path, stream flag, and Codex beta feature header to distinguish the +// native remote compaction v2 wire from the legacy /responses/compact bridge. func HasCompactionTriggerInInput(body []byte) bool { if len(body) == 0 { return false diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go index 70ef3fc01a..70e6af9454 100644 --- a/backend/internal/service/openai_compact_sse_keepalive.go +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -1,6 +1,9 @@ package service import ( + "bufio" + "errors" + "net" "net/http" "sync" "time" @@ -43,12 +46,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) { return func() {} } + originalWriter := c.Writer k := &openAICompactSSEKeepalive{ - writer: c.Writer, + writer: originalWriter, stop: make(chan struct{}), } c.Set(openAICompactSSEKeepaliveKey, k) - c.Writer = &openAICompactKeepaliveWriter{ResponseWriter: c.Writer, k: k} + wrappedWriter := &openAICompactKeepaliveWriter{ResponseWriter: originalWriter, k: k} + c.Writer = wrappedWriter var reqDone <-chan struct{} if c.Request != nil { @@ -71,7 +76,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func timer.Reset(interval) } }() - return k.Stop + return func() { + k.Stop() + // Do not leave a pooled middleware writer reachable through the compact + // wrapper after the request finishes. + if current, ok := c.Writer.(*openAICompactKeepaliveWriter); ok && current == wrappedWriter { + c.Writer = originalWriter + } + } } // beat 在锁内提交(首次)响应头并写出一条 SSE 注释行;返回 false 表示心跳已 @@ -181,52 +193,105 @@ type openAICompactKeepaliveWriter struct { // suspend 停拍心跳;幂等。任何响应构造(含 Header 访问——写响应必先操作 // 响应头)都视为请求侧接管 ResponseWriter。 func (w *openAICompactKeepaliveWriter) suspend() { + if w.k == nil { + return + } w.k.Stop() } func (w *openAICompactKeepaliveWriter) Header() http.Header { w.suspend() + if w.ResponseWriter == nil { + return http.Header{} + } return w.ResponseWriter.Header() } func (w *openAICompactKeepaliveWriter) Write(data []byte) (int, error) { w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } return w.ResponseWriter.Write(data) } func (w *openAICompactKeepaliveWriter) WriteString(s string) (int, error) { w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } return w.ResponseWriter.WriteString(s) } func (w *openAICompactKeepaliveWriter) WriteHeader(code int) { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.WriteHeader(code) } func (w *openAICompactKeepaliveWriter) WriteHeaderNow() { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.WriteHeaderNow() } func (w *openAICompactKeepaliveWriter) Flush() { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.Flush() } +func (w *openAICompactKeepaliveWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + if w.ResponseWriter == nil { + return nil, nil, errors.New("response writer released") + } + return w.ResponseWriter.Hijack() +} + +func (w *openAICompactKeepaliveWriter) CloseNotify() <-chan bool { + if w.ResponseWriter == nil { + ch := make(chan bool) + close(ch) + return ch + } + return w.ResponseWriter.CloseNotify() +} + +func (w *openAICompactKeepaliveWriter) Pusher() http.Pusher { + if w.ResponseWriter == nil { + return nil + } + return w.ResponseWriter.Pusher() +} + func (w *openAICompactKeepaliveWriter) Status() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Status() } func (w *openAICompactKeepaliveWriter) Size() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Size() } func (w *openAICompactKeepaliveWriter) Written() bool { + if w.k == nil || w.ResponseWriter == nil { + return false + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Written() diff --git a/backend/internal/service/openai_compact_sse_keepalive_test.go b/backend/internal/service/openai_compact_sse_keepalive_test.go index 3b217a0718..1efed7e9e4 100644 --- a/backend/internal/service/openai_compact_sse_keepalive_test.go +++ b/backend/internal/service/openai_compact_sse_keepalive_test.go @@ -6,6 +6,7 @@ import ( "testing" "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -141,6 +142,110 @@ func TestOpenAICompactKeepaliveWriter_RequestSideWriteSuspendsBeats(t *testing.T require.Contains(t, rec.Body.String(), `{"error":"local reject"}`) } +func TestOpenAICompactKeepaliveWriter_NilInnerWriter_NoPanic(t *testing.T) { + w := &openAICompactKeepaliveWriter{ + k: &openAICompactSSEKeepalive{stop: make(chan struct{})}, + } + w.ResponseWriter = nil + + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Status()) + }) + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Size()) + }) + assert.NotPanics(t, func() { + assert.False(t, w.Written()) + }) + assert.NotPanics(t, func() { + assert.NotNil(t, w.Header()) + }) + assert.NotPanics(t, func() { + n, err := w.Write([]byte("test")) + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + n, err := w.WriteString("test") + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + w.WriteHeader(http.StatusOK) + }) + assert.NotPanics(t, func() { + w.WriteHeaderNow() + }) + assert.NotPanics(t, func() { + w.Flush() + }) + assert.NotPanics(t, func() { + conn, rw, err := w.Hijack() + assert.Nil(t, conn) + assert.Nil(t, rw) + assert.Error(t, err) + }) + assert.NotPanics(t, func() { + ch := w.CloseNotify() + assert.NotNil(t, ch) + }) + assert.NotPanics(t, func() { + assert.Nil(t, w.Pusher()) + }) +} + +func TestOpenAICompactKeepaliveWriter_NilKeepalive_NoPanic(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + w := &openAICompactKeepaliveWriter{ResponseWriter: c.Writer} + + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Status()) + }) + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Size()) + }) + assert.NotPanics(t, func() { + assert.False(t, w.Written()) + }) + assert.NotPanics(t, func() { + w.Header().Set("X-Test", "ok") + }) + assert.NotPanics(t, func() { + w.WriteHeader(http.StatusAccepted) + }) + assert.NotPanics(t, func() { + n, err := w.WriteString("ok") + assert.Equal(t, 2, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + w.Flush() + }) + require.Equal(t, "ok", rec.Header().Get("X-Test")) + require.Equal(t, "ok", rec.Body.String()) +} + +func TestOpenAICompactKeepaliveWriter_DelegatesWhenReady(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, time.Hour) + defer stop() + + w, ok := c.Writer.(*openAICompactKeepaliveWriter) + require.True(t, ok) + + w.Header().Set("X-Test", "ok") + w.WriteHeader(http.StatusAccepted) + n, err := w.WriteString("ready") + require.NoError(t, err) + require.Equal(t, len("ready"), n) + + require.Equal(t, http.StatusAccepted, w.Status()) + require.Equal(t, len("ready"), w.Size()) + require.True(t, w.Written()) + require.Equal(t, "ok", rec.Header().Get("X-Test")) + require.Equal(t, "ready", rec.Body.String()) +} + // fast policy block 在心跳提交后必须降级为 response.failed 终止事件。 func TestWriteOpenAIFastPolicyBlockedResponse_AfterKeepaliveCommit(t *testing.T) { c, rec := newCompactBridgeTestContext(t, true) diff --git a/backend/internal/service/openai_compat_model_test.go b/backend/internal/service/openai_compat_model_test.go index 7c7ac1b94f..e1007c507a 100644 --- a/backend/internal/service/openai_compat_model_test.go +++ b/backend/internal/service/openai_compat_model_test.go @@ -124,6 +124,55 @@ func TestApplyOpenAICompatModelNormalization(t *testing.T) { }) } +func TestForwardAsAnthropic_UsesExactFableMessagesDispatchModel(t *testing.T) { + t.Parallel() + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{"model":"claude-fable-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_fable","object":"response","model":"gpt-5.6-sol","status":"completed","output":[{"type":"message","id":"msg_fable","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_fable"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}}, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.6-sol") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "claude-fable-5", result.Model) + require.Equal(t, "gpt-5.6-sol", result.BillingModel) + require.Equal(t, "gpt-5.6-sol", result.UpstreamModel) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.NotContains(t, string(upstream.lastBody), "claude-fable-5") + require.Equal(t, "claude-fable-5", gjson.GetBytes(rec.Body.Bytes(), "model").String()) +} + func TestForwardAsAnthropic_NormalizesRoutingAndEffortForGpt54XHigh(t *testing.T) { t.Parallel() gin.SetMode(gin.TestMode) @@ -837,8 +886,7 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) { require.NoError(t, err) require.NotNil(t, firstResult) require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state")) - require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[0].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`) secondRec := httptest.NewRecorder() @@ -852,12 +900,73 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) { require.Equal(t, "turn_state_first", upstream.requests[1].Header.Get("x-codex-turn-state")) require.Equal(t, generateSessionUUID(isolateOpenAISessionID(0, "stable-cache-key")), upstream.requests[1].Header.Get("session_id")) require.Empty(t, upstream.requests[1].Header.Get("conversation_id")) - require.Empty(t, upstream.requests[1].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[1].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists()) require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists()) } +func TestForwardAsAnthropic_OAuthRestoresCodexIdentityHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + + const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)" + tests := []struct { + name string + userAgent string + originator string + wantUserAgent string + wantOriginator string + }{ + { + name: "官方UA逐字保留并重新配对", + userAgent: tuiUA, + originator: "opencode", + wantUserAgent: tuiUA, + wantOriginator: "codex-tui", + }, + { + name: "第三方UA回退为默认Codex身份", + userAgent: "third-party-client/1.0.0", + originator: "opencode", + wantUserAgent: codexCLIUserAgent, + wantOriginator: "codex_cli_rs", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", tt.userAgent) + c.Request.Header.Set("originator", tt.originator) + + upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_identity", "gpt-5.4")} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}}, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.4") + require.NoError(t, err) + require.NotNil(t, result) + requireOpenAIMessagesCodexIdentity(t, upstream.lastReq, tt.wantUserAgent, tt.wantOriginator) + }) + } +} + func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey(t *testing.T) { t.Parallel() gin.SetMode(gin.TestMode) @@ -896,6 +1005,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey firstSessionID := upstream.requests[0].Header.Get("session_id") require.NotEmpty(t, firstSessionID) require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").Exists()) secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`) @@ -910,6 +1020,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey require.Equal(t, firstSessionID, upstream.requests[1].Header.Get("session_id")) require.Equal(t, "turn_state_digest_first", upstream.requests[1].Header.Get("x-codex-turn-state")) require.Empty(t, upstream.requests[1].Header.Get("conversation_id")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists()) require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists()) } @@ -1064,8 +1175,7 @@ func TestForwardAsAnthropic_OAuthKeepsSystemAsDeveloperInput(t *testing.T) { instructions := gjson.GetBytes(upstream.lastBody, "instructions") require.True(t, instructions.Exists()) require.Empty(t, instructions.String()) - require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[0].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") } func TestForwardAsAnthropic_OAuthAddsClaudeCodeTodoGuardForCompatModel(t *testing.T) { @@ -1202,6 +1312,15 @@ func openAICompatSSECompletedResponse(responseID, model string) *http.Response { } } +func requireOpenAIMessagesCodexIdentity(t *testing.T, req *http.Request, wantUserAgent, wantOriginator string) { + t.Helper() + require.NotNil(t, req) + require.Equal(t, wantUserAgent, req.Header.Get("User-Agent")) + require.Equal(t, wantOriginator, req.Header.Get("originator")) + require.Equal(t, codexCLIVersion, req.Header.Get("version")) + require.Equal(t, "responses=experimental", req.Header.Get("OpenAI-Beta")) +} + func openAICompatSSEResponseWithoutUsage(responseID, model string) *http.Response { body := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"` + responseID + `","object":"response","model":"` + model + `","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}]}}`, diff --git a/backend/internal/service/openai_content_session_seed.go b/backend/internal/service/openai_content_session_seed.go index 7c2ba25140..fce85f11bd 100644 --- a/backend/internal/service/openai_content_session_seed.go +++ b/backend/internal/service/openai_content_session_seed.go @@ -11,6 +11,10 @@ import ( // and explicit session IDs (e.g. "sess-xxx" or "compat_cc_xxx"). const contentSessionSeedPrefix = "compat_cs_" +// contentStablePrefixSessionSeedPrefix distinguishes cache identities derived +// only from request fields that remain stable across independent prompts. +const contentStablePrefixSessionSeedPrefix = "compat_csp_" + // deriveOpenAIContentSessionSeed builds a stable session seed from an // OpenAI-format request body. Only fields constant across conversation turns // are included: model, tools/functions definitions, system/developer prompts, @@ -105,3 +109,156 @@ func deriveOpenAIContentSessionSeed(body []byte) string { } return contentSessionSeedPrefix + b.String() } + +// deriveOpenAIAnchoredContentSessionSeed returns the legacy content-derived +// seed only when it contains a meaningful user/input anchor. This preserves +// the existing session derivation while preventing model-only requests from +// becoming a tenant-wide cache routing identity. +func deriveOpenAIAnchoredContentSessionSeed(body []byte) string { + if !hasOpenAIContentSessionUserAnchor(body) { + return "" + } + return deriveOpenAIContentSessionSeed(body) +} + +func hasOpenAIContentSessionUserAnchor(body []byte) bool { + if len(body) == 0 { + return false + } + + if messages := gjson.GetBytes(body, "messages"); messages.Exists() && messages.IsArray() { + anchored := false + messages.ForEach(func(_, message gjson.Result) bool { + if strings.TrimSpace(message.Get("role").String()) != "user" { + return true + } + anchored = hasMeaningfulOpenAIContent(message.Get("content")) + return false + }) + return anchored + } + + input := gjson.GetBytes(body, "input") + if !input.Exists() { + return false + } + if input.Type == gjson.String { + return strings.TrimSpace(input.String()) != "" + } + if !input.IsArray() { + return false + } + + anchored := false + input.ForEach(func(_, item gjson.Result) bool { + if strings.TrimSpace(item.Get("role").String()) == "user" { + anchored = hasMeaningfulOpenAIContent(item.Get("content")) + return false + } + if strings.TrimSpace(item.Get("type").String()) == "input_text" { + anchored = strings.TrimSpace(item.Get("text").String()) != "" + return false + } + return true + }) + return anchored +} + +func hasMeaningfulOpenAIContent(content gjson.Result) bool { + if !content.Exists() || content.Type == gjson.Null { + return false + } + if content.Type == gjson.String { + return strings.TrimSpace(content.String()) != "" + } + if !content.IsArray() { + normalized, ok := normalizeNonEmptyCompatSeedJSON(content) + return ok && strings.TrimSpace(normalized) != "" + } + + meaningful := false + content.ForEach(func(_, item gjson.Result) bool { + if item.Type == gjson.String { + meaningful = strings.TrimSpace(item.String()) != "" + } else if text := item.Get("text"); text.Exists() { + meaningful = strings.TrimSpace(text.String()) != "" + } else { + _, meaningful = normalizeNonEmptyCompatSeedJSON(item) + } + return !meaningful + }) + return meaningful +} + +// deriveOpenAIStablePrefixSessionSeed builds a seed from the reusable prefix +// of an OpenAI-format request. User and assistant content are deliberately +// excluded so independent prompts with the same system/tool prefix can share +// an upstream prompt-cache routing identity. +// +// An empty result means the request has no meaningful stable prefix. Callers +// must then use a narrower fallback instead of grouping all requests by tenant +// and model alone. +func deriveOpenAIStablePrefixSessionSeed(body []byte) string { + if len(body) == 0 { + return "" + } + + var b strings.Builder + hasStablePrefix := false + appendJSON := func(label string, value gjson.Result) { + normalized, ok := normalizeNonEmptyCompatSeedJSON(value) + if !ok { + return + } + _, _ = b.WriteString("|") + _, _ = b.WriteString(label) + _, _ = b.WriteString("=") + _, _ = b.WriteString(normalized) + hasStablePrefix = true + } + + if tools := gjson.GetBytes(body, "tools"); tools.Exists() && tools.IsArray() { + appendJSON("tools", tools) + } + if funcs := gjson.GetBytes(body, "functions"); funcs.Exists() && funcs.IsArray() { + appendJSON("functions", funcs) + } + if instructions := gjson.GetBytes(body, "instructions"); strings.TrimSpace(instructions.String()) != "" { + appendJSON("instructions", instructions) + } + + appendSystemMessages := func(items gjson.Result) { + items.ForEach(func(_, item gjson.Result) bool { + role := strings.TrimSpace(item.Get("role").String()) + switch role { + case "system", "developer": + appendJSON(role, item.Get("content")) + } + return true + }) + } + + if messages := gjson.GetBytes(body, "messages"); messages.Exists() && messages.IsArray() { + appendSystemMessages(messages) + } else if input := gjson.GetBytes(body, "input"); input.Exists() && input.IsArray() { + appendSystemMessages(input) + } + + if !hasStablePrefix { + return "" + } + return contentStablePrefixSessionSeedPrefix + b.String() +} + +func normalizeNonEmptyCompatSeedJSON(value gjson.Result) (string, bool) { + if !value.Exists() || value.Type == gjson.Null { + return "", false + } + normalized := normalizeCompatSeedJSON(json.RawMessage(value.Raw)) + switch normalized { + case "", `""`, "[]", "{}", "null": + return "", false + default: + return normalized, true + } +} diff --git a/backend/internal/service/openai_content_session_seed_test.go b/backend/internal/service/openai_content_session_seed_test.go index 65a0bf1808..6dadc5cf53 100644 --- a/backend/internal/service/openai_content_session_seed_test.go +++ b/backend/internal/service/openai_content_session_seed_test.go @@ -216,3 +216,154 @@ func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_TypedMessageItem(t *testing require.Contains(t, seed, "|first_user=") require.Contains(t, seed, "Hello from typed message") } + +func TestDeriveOpenAIStablePrefixSessionSeed_IgnoresUserContent(t *testing.T) { + first := []byte(`{ + "model": "grok", + "instructions": "Be concise.", + "tools": [{"type":"function","name":"lookup","parameters":{"type":"object"}}], + "input": [{"role":"user","content":"Question A"}] + }`) + second := []byte(`{ + "model": "grok", + "instructions": "Be concise.", + "tools": [{"parameters":{"type":"object"},"name":"lookup","type":"function"}], + "input": [{"role":"user","content":"Question B"}] + }`) + + firstSeed := deriveOpenAIStablePrefixSessionSeed(first) + secondSeed := deriveOpenAIStablePrefixSessionSeed(second) + + require.NotEmpty(t, firstSeed) + require.Equal(t, firstSeed, secondSeed) + require.NotContains(t, firstSeed, "Question A") + require.NotContains(t, firstSeed, "first_user") +} + +func TestDeriveOpenAIStablePrefixSessionSeed_IsolatesStablePrefixFields(t *testing.T) { + base := []byte(`{ + "instructions":"Be concise.", + "tools":[{"type":"function","name":"lookup"}], + "input":[{"role":"system","content":"System A"},{"role":"user","content":"Question"}] + }`) + differentInstructions := []byte(`{ + "instructions":"Be detailed.", + "tools":[{"type":"function","name":"lookup"}], + "input":[{"role":"system","content":"System A"},{"role":"user","content":"Question"}] + }`) + differentTools := []byte(`{ + "instructions":"Be concise.", + "tools":[{"type":"function","name":"search"}], + "input":[{"role":"system","content":"System A"},{"role":"user","content":"Question"}] + }`) + differentSystem := []byte(`{ + "instructions":"Be concise.", + "tools":[{"type":"function","name":"lookup"}], + "input":[{"role":"system","content":"System B"},{"role":"user","content":"Question"}] + }`) + + baseSeed := deriveOpenAIStablePrefixSessionSeed(base) + require.NotEqual(t, baseSeed, deriveOpenAIStablePrefixSessionSeed(differentInstructions)) + require.NotEqual(t, baseSeed, deriveOpenAIStablePrefixSessionSeed(differentTools)) + require.NotEqual(t, baseSeed, deriveOpenAIStablePrefixSessionSeed(differentSystem)) +} + +func TestDeriveOpenAIStablePrefixSessionSeed_ChatSystemAndDeveloper(t *testing.T) { + first := []byte(`{ + "messages":[ + {"role":"system","content":"System prompt"}, + {"role":"developer","content":[{"type":"text","text":"Developer prompt"}]}, + {"role":"user","content":"Question A"} + ] + }`) + second := []byte(`{ + "messages":[ + {"role":"system","content":"System prompt"}, + {"role":"developer","content":[{"text":"Developer prompt","type":"text"}]}, + {"role":"user","content":"Question B"} + ] + }`) + + firstSeed := deriveOpenAIStablePrefixSessionSeed(first) + require.Equal(t, firstSeed, deriveOpenAIStablePrefixSessionSeed(second)) + require.Contains(t, firstSeed, "System prompt") + require.Contains(t, firstSeed, "Developer prompt") +} + +func TestDeriveOpenAIStablePrefixSessionSeed_EncodesSystemAndDeveloperRoles(t *testing.T) { + systemThenDeveloper := []byte(`{ + "messages":[ + {"role":"system","content":"Prompt A"}, + {"role":"developer","content":"Prompt B"} + ] + }`) + developerThenSystem := []byte(`{ + "messages":[ + {"role":"developer","content":"Prompt A"}, + {"role":"system","content":"Prompt B"} + ] + }`) + + firstSeed := deriveOpenAIStablePrefixSessionSeed(systemThenDeveloper) + secondSeed := deriveOpenAIStablePrefixSessionSeed(developerThenSystem) + + require.NotEqual(t, firstSeed, secondSeed) + require.Contains(t, firstSeed, "|system=") + require.Contains(t, firstSeed, "|developer=") +} + +func TestDeriveOpenAIStablePrefixSessionSeed_EncodesInstructionDelimiters(t *testing.T) { + instructionOnly := []byte(`{ + "instructions":"foo|system=\"bar\"" + }`) + instructionAndSystem := []byte(`{ + "instructions":"foo", + "input":[{"role":"system","content":"bar"}] + }`) + + firstSeed := deriveOpenAIStablePrefixSessionSeed(instructionOnly) + secondSeed := deriveOpenAIStablePrefixSessionSeed(instructionAndSystem) + + require.NotEmpty(t, firstSeed) + require.NotEmpty(t, secondSeed) + require.NotEqual(t, firstSeed, secondSeed) +} + +func TestDeriveOpenAIAnchoredContentSessionSeed_RequiresMeaningfulAnchor(t *testing.T) { + emptyAnchors := [][]byte{ + nil, + []byte(`{"model":"grok"}`), + []byte(`{"model":"grok","messages":[{"role":"assistant","content":"answer"}]}`), + []byte(`{"model":"grok","messages":[{"role":"user","content":" "}]}`), + []byte(`{"model":"grok","messages":[{"role":"user","content":[{"type":"text","text":""}]}]}`), + []byte(`{"model":"grok","input":" "}`), + []byte(`{"model":"grok","input":[{"type":"input_text","text":""}]}`), + } + for _, body := range emptyAnchors { + require.Empty(t, deriveOpenAIAnchoredContentSessionSeed(body)) + } + + meaningfulAnchors := [][]byte{ + []byte(`{"model":"grok","messages":[{"role":"user","content":"question"}]}`), + []byte(`{"model":"grok","messages":[{"role":"user","content":[{"type":"text","text":"question"}]}]}`), + []byte(`{"model":"grok","input":"question"}`), + []byte(`{"model":"grok","input":[{"type":"input_text","text":"question"}]}`), + } + for _, body := range meaningfulAnchors { + require.NotEmpty(t, deriveOpenAIAnchoredContentSessionSeed(body)) + } +} + +func TestDeriveOpenAIStablePrefixSessionSeed_RequiresMeaningfulPrefix(t *testing.T) { + tests := [][]byte{ + nil, + []byte(`{}`), + []byte(`{"model":"grok","input":"Question A"}`), + []byte(`{"model":"grok","tools":[],"input":"Question A"}`), + []byte(`{"model":"grok","functions":[],"instructions":" ","messages":[{"role":"system","content":""},{"role":"user","content":"Question A"}]}`), + } + + for _, body := range tests { + require.Empty(t, deriveOpenAIStablePrefixSessionSeed(body)) + } +} diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go index 816f5a26e4..690b9991bb 100644 --- a/backend/internal/service/openai_gateway_cc_pipeline.go +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -88,6 +88,9 @@ func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError( upstreamMsg string, upstreamModel string, ) *UpstreamFailoverError { + if account != nil && account.Platform == PlatformGrok { + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + } if !s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { return nil } @@ -109,7 +112,9 @@ func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError( Message: upstreamMsg, Detail: upstreamDetail, }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) + if account.Platform != PlatformGrok { + s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) + } return &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, @@ -159,6 +164,7 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest( stream bool, bearerToken string, userAgent string, + grokCacheIdentity string, ) (*http.Response, error) { upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(body)) @@ -190,6 +196,10 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest( // 账号级请求头覆写(仅 openai api_key 账号启用时生效) account.ApplyHeaderOverrides(upstreamReq.Header) + if account.Platform == PlatformGrok { + applyGrokCLIHeaders(upstreamReq.Header) + applyGrokCacheHeaders(upstreamReq.Header, grokCacheIdentity) + } proxyURL := "" if account.Proxy != nil { diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 5d80292c47..fe24b3db36 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -72,6 +72,16 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( } if account.Platform == PlatformGrok { + if account.IsGrokOAuth() { + if eligible, reason := grokChatResponsesBridgeEligibility(body); eligible { + return s.forwardGrokChatCompletionsViaResponses(ctx, c, account, body, promptCacheKey, defaultMappedModel) + } else { + logger.L().Debug("grok chat_completions: using raw fallback", + zap.Int64("account_id", account.ID), + zap.String("reason", reason), + ) + } + } return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) } diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 6def67c7ec..7693986b0d 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -76,6 +76,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( // 2. Resolve model mapping (same as ForwardAsChatCompletions) billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + grokCacheIdentity := "" + if account.Platform == PlatformGrok { + // Resolve before image bridging or other body rewrites so the fallback is + // anchored to the client's stable conversation prefix. + grokCacheIdentity = resolveGrokCacheIdentity(c, body, "", upstreamModel) + } reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) // 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) @@ -134,6 +140,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( return nil, fmt.Errorf("enable stream usage: %w", usageErr) } } + if account.Platform == PlatformGrok { + upstreamBody, err = stripGrokChatPromptCacheKey(upstreamBody) + if err != nil { + return nil, fmt.Errorf("remove Responses-only Grok prompt cache key: %w", err) + } + } logger.L().Debug("openai chat_completions raw: forwarding without protocol conversion", zap.Int64("account_id", account.ID), @@ -148,11 +160,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( if err != nil { return nil, err } + SetActualOpenAIUpstreamEndpoint(c, grokChatRawEndpoint) customUA := account.GetOpenAIUserAgent() if customUA == "" && account.Platform == PlatformGrok { customUA = "sub2api-grok/1.0" } - resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA) + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA, grokCacheIdentity) if err != nil { return nil, err } @@ -162,7 +175,6 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( if resp.StatusCode >= 400 { respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ Platform: account.Platform, AccountID: account.ID, @@ -189,7 +201,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) } // 8. Forward response @@ -202,6 +214,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } if result != nil { addOpenAIUsage(&result.Usage, bridgeUsage) + result.UpstreamEndpoint = grokChatRawEndpoint } return result, forwardErr } diff --git a/backend/internal/service/openai_gateway_chat_completions_test.go b/backend/internal/service/openai_gateway_chat_completions_test.go index b85ee33947..5186598a70 100644 --- a/backend/internal/service/openai_gateway_chat_completions_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_test.go @@ -98,7 +98,7 @@ func TestNormalizeResponsesBodyServiceTier(t *testing.T) { require.False(t, gjson.GetBytes(body, "service_tier").Exists()) } -func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *testing.T) { +func TestForwardAsChatCompletions_UnknownModelWithoutMessagesDispatchKeepsRequestedModel(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() @@ -129,7 +129,7 @@ func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *te }, } - result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.4") + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") require.Error(t, err) require.Nil(t, result) require.Equal(t, "gpt6", gjson.GetBytes(upstream.lastBody, "model").String()) diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index e1c9aea15f..509eab8d71 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -42,6 +42,20 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if normalized { body = normalizedBody } + wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account) + // 仅允许 WS 入站请求走 WS 上游,避免出现 HTTP -> WS 协议混用。 + wsDecision = resolveOpenAIWSDecisionByClientTransport(wsDecision, GetOpenAIClientTransport(c)) + passthroughEnabled := account.IsOpenAIPassthroughEnabled() + if shouldFlattenOpenAIResponsesNamespaces(account, wsDecision.Transport, passthroughEnabled) { + body, err = flattenOpenAIResponsesNamespaces(c, body) + if err != nil { + setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "") + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", "message": err.Error(), "param": "tools", + }}) + return nil, err + } + } originalBody := body requestView := newOpenAIRequestView(body) @@ -49,7 +63,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco originalModel := reqModel if account.Platform == PlatformGrok { - _ = promptCacheKey return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime) } @@ -65,10 +78,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if isCodexCLI { codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() } - wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account) - clientTransport := GetOpenAIClientTransport(c) - // 仅允许 WS 入站请求走 WS 上游,避免出现 HTTP -> WS 协议混用。 - wsDecision = resolveOpenAIWSDecisionByClientTransport(wsDecision, clientTransport) if c != nil { c.Set("openai_ws_transport_decision", string(wsDecision.Transport)) c.Set("openai_ws_transport_reason", wsDecision.Reason) @@ -97,7 +106,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } return nil, errors.New("openai ws v1 is temporarily unsupported; use ws v2") } - passthroughEnabled := account.IsOpenAIPassthroughEnabled() if passthroughEnabled { if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { strippedBody, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(body) @@ -174,7 +182,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if apiKey != nil { imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group) } - codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) + codexImageGenerationBridgeEnabled := isCodexCLI && + !isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader)) && + imageGenerationAllowed && + codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && + s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) var imageIntent bool if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { decoded, decodeErr := ensureReqBody() diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index ef823fa2fd..8bd1d456dc 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io" + "log/slog" "net/http" "strings" "time" @@ -20,6 +21,10 @@ import ( const ( grokComposerImageBridgeVisionModel = "grok-build-0.1" grokComposerImageBridgeMaxOutputTokens = 512 + grokUpstreamUserAgent = "sub2api-grok/1.0" + grokCLIVersion = "0.2.93" + grokDefaultResponsesModel = "grok-4.5" + grokRateLimitFallbackCooldown = 2 * time.Minute ) func (s *OpenAIGatewayService) forwardGrokResponses( @@ -31,18 +36,23 @@ func (s *OpenAIGatewayService) forwardGrokResponses( reqStream bool, startTime time.Time, ) (*OpenAIForwardResult, error) { - if account.Type != AccountTypeOAuth { - return nil, fmt.Errorf("grok account type %s is not supported by subscription forwarding", account.Type) + if account.Type != AccountTypeOAuth && account.Type != AccountTypeAPIKey { + return nil, fmt.Errorf("grok account type %s is not supported by Responses forwarding", account.Type) } upstreamModel := account.GetMappedModel(originalModel) if strings.TrimSpace(upstreamModel) == "" { - upstreamModel = "grok-4.3" + upstreamModel = grokDefaultResponsesModel } + cacheIdentity := resolveGrokCacheIdentity(c, body, "", upstreamModel) patchedBody, err := patchGrokResponsesBody(body, upstreamModel) if err != nil { return nil, err } + patchedBody, err = applyGrokResponsesCacheIdentity(patchedBody, body, cacheIdentity, account.IsGrokOAuth()) + if err != nil { + return nil, fmt.Errorf("apply grok prompt cache identity: %w", err) + } token, _, err := s.GetAccessToken(ctx, account) if err != nil { @@ -51,7 +61,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) defer releaseUpstreamCtx() - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity) if err != nil { return nil, err } @@ -72,7 +82,6 @@ func (s *OpenAIGatewayService) forwardGrokResponses( if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) resp.Body = io.NopCloser(bytes.NewReader(respBody)) - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) upstreamMsg := sanitizeUpstreamErrorMessage(extractUpstreamErrorMessage(respBody)) if upstreamMsg == "" { upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode) @@ -97,7 +106,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses( return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel) } - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) var usage *OpenAIUsage var firstTokenMs *int @@ -146,6 +155,10 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { if err != nil { return nil, err } + out, err = sanitizeGrokResponsesModelCapabilities(out, upstreamModel) + if err != nil { + return nil, err + } for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} { if gjson.GetBytes(out, unsupportedField).Exists() { out, err = sjson.DeleteBytes(out, unsupportedField) @@ -168,6 +181,14 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { if err != nil { return nil, err } + out, err = sanitizeGrokResponsesInput(out) + if err != nil { + return nil, err + } + out, err = sanitizeGrokReasoningNullContent(out) + if err != nil { + return nil, err + } out, err = sanitizeGrokResponsesTools(out) if err != nil { return nil, err @@ -175,6 +196,38 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { return out, nil } +func sanitizeGrokResponsesModelCapabilities(body []byte, upstreamModel string) ([]byte, error) { + if !grokModelRejectsReasoningEffort(upstreamModel) { + return body, nil + } + + out := body + for _, field := range []string{"reasoning", "reasoning_effort", "reasoningEffort"} { + if !gjson.GetBytes(out, field).Exists() { + continue + } + var err error + out, err = sjson.DeleteBytes(out, field) + if err != nil { + return nil, fmt.Errorf("remove unsupported Grok Composer %s: %w", field, err) + } + } + return out, nil +} + +func grokModelRejectsReasoningEffort(model string) bool { + model = strings.TrimSpace(strings.ToLower(model)) + if slash := strings.LastIndex(model, "/"); slash >= 0 { + model = strings.TrimSpace(model[slash+1:]) + } + switch model { + case "grok-composer", "grok-composer-2.5-fast", "composer-2.5": + return true + default: + return false + } +} + var grokResponsesUnsupportedRecursiveFields = map[string]struct{}{ "external_web_access": {}, } @@ -223,6 +276,67 @@ func deleteJSONFields(value any, fields map[string]struct{}) bool { } } +// additional_tools is a Codex/Responses Lite private input carrier. xAI's +// Responses schema accepts ordinary message/function-call input items but +// rejects this carrier before inference with a ModelInput deserialization +// error. Top-level supported tools remain available through the separate +// sanitizeGrokResponsesTools path. +func sanitizeGrokResponsesInput(body []byte) ([]byte, error) { + if !bytes.Contains(body, []byte(`"additional_tools"`)) { + return body, nil + } + input := gjson.GetBytes(body, "input") + if !input.Exists() || !input.IsArray() { + return body, nil + } + + rawItems := input.Array() + filtered := make([]json.RawMessage, 0, len(rawItems)) + for _, item := range rawItems { + if strings.TrimSpace(item.Get("type").String()) == "additional_tools" { + continue + } + filtered = append(filtered, json.RawMessage(item.Raw)) + } + if len(filtered) == len(rawItems) { + return body, nil + } + encoded, err := json.Marshal(filtered) + if err != nil { + return nil, err + } + return sjson.SetRawBytes(body, "input", encoded) +} + +// sanitizeGrokReasoningNullContent 删除 reasoning 项中的 "content": null。 +// xAI 的 untagged enum 反序列化器拒收该字段,返回 422。 +func sanitizeGrokReasoningNullContent(body []byte) ([]byte, error) { + input := gjson.GetBytes(body, "input") + if !input.Exists() || !input.IsArray() { + return body, nil + } + + items := input.Array() + changed := false + for i := len(items) - 1; i >= 0; i-- { + item := items[i] + if strings.TrimSpace(item.Get("type").String()) != "reasoning" { + continue + } + contentResult := item.Get("content") + if contentResult.Exists() && contentResult.Type == gjson.Null { + var err error + body, err = sjson.DeleteBytes(body, fmt.Sprintf("input.%d.content", i)) + if err != nil { + return nil, err + } + changed = true + } + } + _ = changed + return body, nil +} + var grokResponsesSupportedToolTypes = map[string]struct{}{ "code_execution": {}, "code_interpreter": {}, @@ -457,7 +571,9 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token) + // Image-description probes are auxiliary requests, not conversation turns. + // Do not bind them to the caller's Grok prompt-cache identity. + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "") releaseUpstreamCtx() if err != nil { return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err) @@ -476,7 +592,6 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) upstreamMsg := sanitizeUpstreamErrorMessage(extractUpstreamErrorMessage(respBody)) if upstreamMsg == "" { upstreamMsg = fmt.Sprintf("xAI image bridge upstream returned status %d", resp.StatusCode) @@ -501,7 +616,7 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( return "", OpenAIUsage{}, fmt.Errorf("grok composer image bridge upstream error: %s", upstreamMsg) } - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, nil) if err != nil { return "", OpenAIUsage{}, fmt.Errorf("read grok composer image bridge response: %w", err) @@ -623,7 +738,7 @@ func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) { dst.ImageOutputTokens += usage.ImageOutputTokens } -func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { +func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string) (*http.Request, error) { targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) if err != nil { return nil, err @@ -635,7 +750,8 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc req.Header.Set("Authorization", "Bearer "+token) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json, text/event-stream") - req.Header.Set("User-Agent", "sub2api-grok/1.0") + applyGrokCLIHeaders(req.Header) + applyGrokCacheHeaders(req.Header, cacheIdentity) if c != nil { if v := c.GetHeader("OpenAI-Beta"); strings.TrimSpace(v) != "" { req.Header.Set("OpenAI-Beta", v) @@ -644,33 +760,201 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc return req, nil } -func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, accountID int64, snapshot *xai.QuotaSnapshot) { - if s == nil || s.accountRepo == nil || accountID <= 0 || snapshot == nil { +// applyGrokCLIHeaders identifies subscription traffic as a supported Grok CLI +// version. The CLI gateway rejects otherwise valid OAuth requests without it. +func applyGrokCLIHeaders(headers http.Header) { + if headers == nil { return } - if s.codexSnapshotThrottle != nil && !s.codexSnapshotThrottle.Allow(accountID, time.Now()) { + headers.Set("User-Agent", grokUpstreamUserAgent) + headers.Set("X-Grok-Client-Version", grokCLIVersion) +} + +func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) { + if s == nil || account == nil || account.ID <= 0 || snapshot == nil { return } - _ = s.accountRepo.UpdateExtra(ctx, accountID, map[string]any{ - grokQuotaSnapshotExtraKey: snapshot, - }) + accountID := account.ID + now := time.Now() + resetAt, hasActiveLimit := grokRateLimitResetAt(snapshot, now) + if hasActiveLimit { + normalizeGrokExhaustedWindowResets(snapshot, resetAt, now) + } + critical := snapshot.StatusCode == http.StatusTooManyRequests || hasActiveLimit + if s.codexSnapshotThrottle != nil { + allowed := s.codexSnapshotThrottle.Allow(accountID, now) + if !critical && !allowed { + return + } + } + + stateCtx := ctx + if hasActiveLimit { + var cancel context.CancelFunc + stateCtx, cancel = openAIAccountStateContext(ctx) + defer cancel() + } + if s.accountRepo != nil { + _ = s.accountRepo.UpdateExtra(stateCtx, accountID, map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + }) + } + // Error responses are reconciled by handleGrokAccountUpstreamError, which + // also installs the immediate in-memory scheduling block. Successful + // responses can still consume the last available request/token, so persist + // that exhausted window here as a real rate limit rather than relying only + // on the passive snapshot scheduler check. + if hasActiveLimit { + s.rateLimitGrok(stateCtx, account, resetAt) + } +} + +func parseGrokQuotaSnapshot(headers http.Header, statusCode int, now time.Time) *xai.QuotaSnapshot { + snapshot := xai.ParseQuotaHeaders(headers, statusCode) + if snapshot == nil && statusCode == http.StatusTooManyRequests { + return &xai.QuotaSnapshot{ + StatusCode: statusCode, + UpdatedAt: now.UTC().Format(time.RFC3339), + } + } + return snapshot +} + +func normalizeGrokExhaustedWindowResets(snapshot *xai.QuotaSnapshot, resetAt, now time.Time) { + if snapshot == nil || !resetAt.After(now) { + return + } + for _, window := range []*xai.QuotaWindow{snapshot.Requests, snapshot.Tokens} { + if window == nil || window.Remaining == nil || *window.Remaining > 0 { + continue + } + candidate := time.Time{} + if window.ResetUnix != nil && *window.ResetUnix > 0 { + candidate = time.Unix(*window.ResetUnix, 0) + } else if parsed, err := time.Parse(time.RFC3339, strings.TrimSpace(window.ResetAt)); err == nil { + candidate = parsed + } + if !candidate.After(now) { + candidate = resetAt + } + resetUnix := candidate.Unix() + window.ResetUnix = &resetUnix + window.ResetAt = candidate.UTC().Format(time.RFC3339) + } +} + +func grokRateLimitResetAt(snapshot *xai.QuotaSnapshot, now time.Time) (time.Time, bool) { + if snapshot == nil { + return time.Time{}, false + } + + // Retry-After is xAI's explicit retry boundary. Use the observation time so + // a persisted snapshot does not start a fresh cooldown every time it is read. + retryAfterExpired := false + var resetAt time.Time + if snapshot.RetryAfterSeconds != nil && *snapshot.RetryAfterSeconds > 0 { + observedAt := now + if parsed, err := time.Parse(time.RFC3339, strings.TrimSpace(snapshot.UpdatedAt)); err == nil { + observedAt = parsed + } + retryAfterResetAt := observedAt.Add(time.Duration(*snapshot.RetryAfterSeconds) * time.Second) + if retryAfterResetAt.After(now) { + resetAt = retryAfterResetAt + } else { + retryAfterExpired = true + } + } + + exhausted := false + for _, window := range []*xai.QuotaWindow{snapshot.Requests, snapshot.Tokens} { + if window == nil || window.Remaining == nil || *window.Remaining > 0 { + continue + } + exhausted = true + candidate := time.Time{} + if window.ResetUnix != nil && *window.ResetUnix > 0 { + candidate = time.Unix(*window.ResetUnix, 0) + } else if parsed, err := time.Parse(time.RFC3339, strings.TrimSpace(window.ResetAt)); err == nil { + candidate = parsed + } + if candidate.After(now) && candidate.After(resetAt) { + resetAt = candidate + } + } + if !resetAt.IsZero() { + return resetAt, true + } + // An observed Retry-After is an absolute boundary once combined with the + // snapshot timestamp. Do not turn an expired persisted snapshot into a new + // rolling fallback cooldown, but still allow a later explicit window reset. + if retryAfterExpired { + return time.Time{}, false + } + if exhausted || snapshot.StatusCode == http.StatusTooManyRequests { + return now.Add(grokRateLimitFallbackCooldown), true + } + return time.Time{}, false +} + +func normalizeGrokRateLimitResetAt(account *Account, resetAt, now time.Time) time.Time { + if !resetAt.After(now) { + resetAt = now.Add(grokRateLimitFallbackCooldown) + } + if account != nil && account.RateLimitResetAt != nil && account.RateLimitResetAt.After(resetAt) { + resetAt = *account.RateLimitResetAt + } + return resetAt +} + +type grokRateLimitExtendingRepository interface { + SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error +} + +func persistGrokRateLimit(ctx context.Context, repo AccountRepository, account *Account, resetAt time.Time) { + if repo == nil || account == nil || account.ID <= 0 { + return + } + resetAt = normalizeGrokRateLimitResetAt(account, resetAt, time.Now()) + stateCtx, cancel := openAIAccountStateContext(ctx) + defer cancel() + var err error + if extendingRepo, ok := repo.(grokRateLimitExtendingRepository); ok { + err = extendingRepo.SetRateLimitedIfLater(stateCtx, account.ID, resetAt) + } else { + err = repo.SetRateLimited(stateCtx, account.ID, resetAt) + } + if err != nil { + slog.Warn("persist_grok_rate_limit_failed", "account_id", account.ID, "reset_at", resetAt.UTC(), "error", err) + } +} + +func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Account, resetAt time.Time) { + if s == nil || account == nil { + return + } + resetAt = normalizeGrokRateLimitResetAt(account, resetAt, time.Now()) + + runtimeUntil := resetAt + if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(runtimeUntil) { + runtimeUntil = *account.TempUnschedulableUntil + } + s.BlockAccountScheduling(account, runtimeUntil, "429") + persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) } func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) { if s == nil || account == nil { return } + now := time.Now() + s.updateGrokUsageSnapshot(ctx, account, parseGrokQuotaSnapshot(headers, statusCode, now)) switch statusCode { case http.StatusUnauthorized: - s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok oauth token unauthorized") + s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok credentials unauthorized") case http.StatusForbidden: - s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok entitlement or subscription tier denied") + s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok access or entitlement denied") case http.StatusTooManyRequests: - cooldown := 2 * time.Minute - if snapshot := xai.ParseQuotaHeaders(headers, statusCode); snapshot != nil && snapshot.RetryAfterSeconds != nil && *snapshot.RetryAfterSeconds > 0 { - cooldown = time.Duration(*snapshot.RetryAfterSeconds) * time.Second - } - s.tempUnscheduleGrok(ctx, account, cooldown, "grok rate limited") + // updateGrokUsageSnapshot installs both runtime and durable rate-limit state. default: if statusCode >= 500 { s.tempUnscheduleGrok(ctx, account, 2*time.Minute, "grok upstream temporary error") diff --git a/backend/internal/service/openai_gateway_grok_cache.go b/backend/internal/service/openai_gateway_grok_cache.go new file mode 100644 index 0000000000..1d689bce8a --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_cache.go @@ -0,0 +1,155 @@ +package service + +import ( + "fmt" + "net/http" + "strings" + + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +const ( + grokConversationIDHeader = "X-Grok-Conv-Id" + grokFreeCacheNativeToolsJSON = `[{"type":"web_search"},{"type":"x_search"}]` + grokFreeCacheDisabledToolChoice = "none" +) + +// resolveGrokCacheIdentity derives one stable, tenant-isolated routing identity +// for xAI's server-side prompt cache. The returned value is safe to expose to +// the upstream: it never contains the client's raw session identifier. +// +// A valid downstream API key is required. This intentionally fails closed on +// internal probes and incomplete request contexts instead of creating a cache +// identity that could be shared by unrelated tenants. +func resolveGrokCacheIdentity(c *gin.Context, body []byte, explicitKey, upstreamModel string) string { + apiKeyID := getAPIKeyIDFromContext(c) + if apiKeyID <= 0 { + return "" + } + // /responses/compact rejects tool_choice and does not represent a normal + // conversation turn. Keep both cache identity and Free-tier routing + // augmentation out of this path. + if isOpenAIResponsesCompactPath(c) { + return "" + } + + model := strings.ToLower(strings.TrimSpace(upstreamModel)) + if model == "" { + return "" + } + + seed := explicitGrokCacheSeed(c, body, explicitKey) + if seed == "" { + seed = deriveOpenAIStablePrefixSessionSeed(body) + if seed == "" { + // A model alone is too broad for cache routing. Preserve the + // existing first-user-derived identity when no reusable prefix is + // available so unrelated prompts do not share one tenant-wide key. + seed = deriveOpenAIAnchoredContentSessionSeed(body) + } + } + if seed == "" { + return "" + } + + // generateSessionUUID hashes the whole seed before formatting it as a UUID. + // Include a versioned namespace so this identity cannot collide with other + // upstream session identifiers derived by sub2api. + isolatedSeed := fmt.Sprintf("grok-prompt-cache:v1:%d:%s:%s", apiKeyID, model, seed) + return generateSessionUUID(isolatedSeed) +} + +func explicitGrokCacheSeed(c *gin.Context, body []byte, explicitKey string) string { + seed := "" + if c != nil { + seed = strings.TrimSpace(c.GetHeader("session_id")) + if seed == "" { + seed = strings.TrimSpace(c.GetHeader("conversation_id")) + } + if seed == "" { + seed = strings.TrimSpace(c.GetHeader(grokConversationIDHeader)) + } + } + if seed == "" && len(body) > 0 { + seed = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) + } + if seed == "" { + seed = strings.TrimSpace(explicitKey) + } + return seed +} + +func isGrokRequestContext(c *gin.Context) bool { + if c == nil { + return false + } + v, exists := c.Get("api_key") + if !exists { + return false + } + apiKey, ok := v.(*APIKey) + return ok && apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == PlatformGrok +} + +// applyGrokResponsesCacheIdentity writes the cache routing identity into an +// xAI Responses request. Existing client values are deliberately replaced by +// the tenant-isolated value to prevent collisions on shared OAuth accounts. +// +// Free OAuth requests without native search tools are routed by xAI to the +// non-cacheable build-free model. For otherwise tool-free requests, add the +// native tools with tool_choice=none: this selects the cache-capable tier +// without allowing an actual search. Any explicit client tools or tool_choice +// disable this augmentation so client function-calling semantics stay intact. +func applyGrokResponsesCacheIdentity(body, intentSourceBody []byte, identity string, injectFreeTierTools bool) ([]byte, error) { + identity = strings.TrimSpace(identity) + if identity == "" { + if gjson.GetBytes(body, "prompt_cache_key").Exists() { + return sjson.DeleteBytes(body, "prompt_cache_key") + } + return body, nil + } + out, err := sjson.SetBytes(body, "prompt_cache_key", identity) + if err != nil { + return nil, err + } + if !injectFreeTierTools { + return out, nil + } + // Inspect the pre-sanitization source. patchGrokResponsesBody may remove an + // unsupported client tool and its tool_choice; that must not turn an + // explicit client tool intent into an eligible native-tool request. + if gjson.GetBytes(intentSourceBody, "tools").Exists() || gjson.GetBytes(intentSourceBody, "tool_choice").Exists() { + return out, nil + } + out, err = sjson.SetRawBytes(out, "tools", []byte(grokFreeCacheNativeToolsJSON)) + if err != nil { + return nil, err + } + return sjson.SetBytes(out, "tool_choice", grokFreeCacheDisabledToolChoice) +} + +// applyGrokCacheHeaders applies the documented Chat Completions conversation +// routing header. The request is built from a fresh header map, so client +// supplied x-grok headers cannot override this server-derived value. +func applyGrokCacheHeaders(headers http.Header, identity string) { + if headers == nil { + return + } + identity = strings.TrimSpace(identity) + if identity == "" { + headers.Del(grokConversationIDHeader) + return + } + headers.Set(grokConversationIDHeader, identity) +} + +// stripGrokChatPromptCacheKey removes the Responses-only body field after it +// has been used as an identity seed. Chat Completions routes cache by header. +func stripGrokChatPromptCacheKey(body []byte) ([]byte, error) { + if !gjson.GetBytes(body, "prompt_cache_key").Exists() { + return body, nil + } + return sjson.DeleteBytes(body, "prompt_cache_key") +} diff --git a/backend/internal/service/openai_gateway_grok_cache_test.go b/backend/internal/service/openai_gateway_grok_cache_test.go new file mode 100644 index 0000000000..42abfc5800 --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_cache_test.go @@ -0,0 +1,330 @@ +//go:build unit + +package service + +import ( + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func newGrokCacheTestContext(apiKeyID int64) *gin.Context { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + if apiKeyID > 0 { + c.Set("api_key", &APIKey{ID: apiKeyID, Group: &Group{Platform: PlatformGrok}}) + } + return c +} + +func TestResolveGrokCacheIdentityStableAcrossAppendOnlyTurns(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(101) + round1 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"}]}`) + round2 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"},{"role":"assistant","content":"first answer"},{"role":"user","content":"second question"}]}`) + + first := resolveGrokCacheIdentity(c, round1, "", "grok-4.5") + second := resolveGrokCacheIdentity(c, round2, "", "grok-4.5") + + require.NotEmpty(t, first) + require.Len(t, first, 36) + require.Equal(t, first, second) +} + +func TestResolveGrokCacheIdentityStableAcrossIndependentPromptsWithSamePrefix(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(102) + firstBody := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"user","content":"Question A"}]}`) + secondBody := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"user","content":"Question B"}]}`) + + first := resolveGrokCacheIdentity(c, firstBody, "", "grok-4.5") + second := resolveGrokCacheIdentity(c, secondBody, "", "grok-4.5") + + require.NotEmpty(t, first) + require.Equal(t, first, second) +} + +func TestResolveGrokCacheIdentityStablePrefixIsolation(t *testing.T) { + gin.SetMode(gin.TestMode) + baseBody := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System A"},{"role":"user","content":"Question A"}]}`) + differentInstructions := []byte(`{"model":"grok","instructions":"be detailed","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System A"},{"role":"user","content":"Question B"}]}`) + differentSystem := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System B"},{"role":"user","content":"Question B"}]}`) + differentTools := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"search"}],"input":[{"role":"system","content":"System A"},{"role":"user","content":"Question B"}]}`) + + base := resolveGrokCacheIdentity(newGrokCacheTestContext(103), baseBody, "", "grok-4.5") + require.NotEqual(t, base, resolveGrokCacheIdentity(newGrokCacheTestContext(104), baseBody, "", "grok-4.5")) + require.NotEqual(t, base, resolveGrokCacheIdentity(newGrokCacheTestContext(103), baseBody, "", "grok-4.3")) + require.NotEqual(t, base, resolveGrokCacheIdentity(newGrokCacheTestContext(103), differentInstructions, "", "grok-4.5")) + require.NotEqual(t, base, resolveGrokCacheIdentity(newGrokCacheTestContext(103), differentSystem, "", "grok-4.5")) + require.NotEqual(t, base, resolveGrokCacheIdentity(newGrokCacheTestContext(103), differentTools, "", "grok-4.5")) +} + +func TestResolveGrokCacheIdentityFallsBackWhenStablePrefixIsEmpty(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(105) + firstBody := []byte(`{"model":"grok","tools":[],"input":"Question A"}`) + secondBody := []byte(`{"model":"grok","tools":[],"input":"Question B"}`) + + first := resolveGrokCacheIdentity(c, firstBody, "", "grok-4.5") + second := resolveGrokCacheIdentity(c, secondBody, "", "grok-4.5") + + require.NotEmpty(t, first) + require.NotEmpty(t, second) + require.NotEqual(t, first, second) +} + +func TestResolveGrokCacheIdentitySkipsUnanchoredFallback(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(106) + tests := [][]byte{ + []byte(`{"model":"grok"}`), + []byte(`{"model":"grok","messages":[{"role":"assistant","content":"answer"}]}`), + []byte(`{"model":"grok","messages":[{"role":"user","content":""}]}`), + []byte(`{"model":"grok","input":" "}`), + } + + for _, body := range tests { + require.Empty(t, resolveGrokCacheIdentity(c, body, "", "grok-4.5")) + } +} + +func TestResolveGrokCacheIdentityIsolatesAPIKeyAndMappedModel(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"grok","input":"same prompt"}`) + + base := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.5") + otherTenant := resolveGrokCacheIdentity(newGrokCacheTestContext(202), body, "", "grok-4.5") + otherModel := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.3") + + require.NotEmpty(t, base) + require.NotEqual(t, base, otherTenant) + require.NotEqual(t, base, otherModel) +} + +func TestResolveGrokCacheIdentityUsesAndIsolatesNativeConversationHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(301) + c.Request.Header.Set(grokConversationIDHeader, "raw-native-conversation") + body1 := []byte(`{"model":"grok","input":"one"}`) + body2 := []byte(`{"model":"grok","input":"different body that must not replace the explicit session"}`) + + first := resolveGrokCacheIdentity(c, body1, "body-cache-key", "grok-4.5") + second := resolveGrokCacheIdentity(c, body2, "another-body-cache-key", "grok-4.5") + + require.Equal(t, "raw-native-conversation", (&OpenAIGatewayService{}).ExtractSessionID(c, body1)) + require.Equal(t, first, second) + require.NotEqual(t, "raw-native-conversation", first) + require.NotContains(t, first, "raw-native-conversation") +} + +func TestResolveGrokCacheIdentityExplicitHeaderPriority(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"grok","prompt_cache_key":"body-key","input":"hi"}`) + c := newGrokCacheTestContext(401) + c.Request.Header.Set(grokConversationIDHeader, "grok-key") + c.Request.Header.Set("conversation_id", "conversation-key") + c.Request.Header.Set("session_id", "session-key") + + got := resolveGrokCacheIdentity(c, body, "explicit-argument", "grok-4.5") + onlySession := newGrokCacheTestContext(401) + onlySession.Request.Header.Set("session_id", "session-key") + want := resolveGrokCacheIdentity(onlySession, []byte(`{"model":"grok","input":"unrelated"}`), "", "grok-4.5") + + require.Equal(t, want, got) +} + +func TestResolveGrokCacheIdentityFailsClosedWithoutAPIKeyContext(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(0) + c.Request.Header.Set(grokConversationIDHeader, "native-session") + + require.Empty(t, resolveGrokCacheIdentity(c, []byte(`{"model":"grok","input":"hi"}`), "", "grok-4.5")) + require.Empty(t, resolveGrokCacheIdentity(nil, []byte(`{"model":"grok","prompt_cache_key":"key"}`), "key", "grok-4.5")) +} + +func TestGrokConversationHeaderIsScopedToGrokRequestScheduling(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"grok","prompt_cache_key":"body-session","input":"hi"}`) + + grokContext := newGrokCacheTestContext(601) + grokContext.Request.Header.Set(grokConversationIDHeader, "native-grok-session") + require.Equal(t, "native-grok-session", (&OpenAIGatewayService{}).ExtractSessionID(grokContext, body)) + + openAIContext := newGrokCacheTestContext(601) + openAIContext.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}}) + openAIContext.Request.Header.Set(grokConversationIDHeader, "must-be-ignored") + require.Equal(t, "body-session", (&OpenAIGatewayService{}).ExtractSessionID(openAIContext, body)) + + withoutGrokHeader := newGrokCacheTestContext(601) + withoutGrokHeader.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}}) + require.Equal(t, + (&OpenAIGatewayService{}).GenerateSessionHash(withoutGrokHeader, body), + (&OpenAIGatewayService{}).GenerateSessionHash(openAIContext, body), + ) +} + +func TestApplyGrokCacheIdentityWritesResponsesBodyAndHeader(t *testing.T) { + sourceBody := []byte(`{"model":"grok-4.5","prompt_cache_key":"raw-client-key"}`) + body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", true) + require.NoError(t, err) + require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String()) + require.Equal(t, "web_search", gjson.GetBytes(body, "tools.0.type").String()) + require.Equal(t, "x_search", gjson.GetBytes(body, "tools.1.type").String()) + require.Equal(t, grokFreeCacheDisabledToolChoice, gjson.GetBytes(body, "tool_choice").String()) + + headers := make(http.Header) + headers.Set(grokConversationIDHeader, "spoofed-client-value") + applyGrokCacheHeaders(headers, "isolated-id") + require.Equal(t, "isolated-id", headers.Get(grokConversationIDHeader)) + applyGrokCacheHeaders(headers, "") + require.Empty(t, headers.Get(grokConversationIDHeader)) + + chatBody, err := stripGrokChatPromptCacheKey(body) + require.NoError(t, err) + require.False(t, gjson.GetBytes(chatBody, "prompt_cache_key").Exists()) + + unscopedSourceBody := []byte(`{"model":"grok","prompt_cache_key":"raw-client-key"}`) + unscopedBody, err := applyGrokResponsesCacheIdentity(unscopedSourceBody, unscopedSourceBody, "", true) + require.NoError(t, err) + require.False(t, gjson.GetBytes(unscopedBody, "prompt_cache_key").Exists()) + require.False(t, gjson.GetBytes(unscopedBody, "tools").Exists()) + require.False(t, gjson.GetBytes(unscopedBody, "tool_choice").Exists()) +} + +func TestApplyGrokCacheIdentityPreservesExplicitClientToolFields(t *testing.T) { + tests := []struct { + name string + body string + }{ + { + name: "tools only", + body: `{"model":"grok","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}]}`, + }, + { + name: "empty tools array", + body: `{"model":"grok","tools":[]}`, + }, + { + name: "null tools", + body: `{"model":"grok","tools":null}`, + }, + { + name: "tool choice only", + body: `{"model":"grok","tool_choice":{"type":"function","name":"lookup"}}`, + }, + { + name: "null tool choice", + body: `{"model":"grok","tool_choice":null}`, + }, + { + name: "both fields", + body: `{"model":"grok","tools":[{"type":"web_search"}],"tool_choice":"auto"}`, + }, + { + name: "unsupported tool", + body: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeTools := gjson.Get(tt.body, "tools") + beforeChoice := gjson.Get(tt.body, "tool_choice") + body, err := applyGrokResponsesCacheIdentity([]byte(tt.body), []byte(tt.body), "isolated-id", true) + require.NoError(t, err) + require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String()) + require.Equal(t, beforeTools.Exists(), gjson.GetBytes(body, "tools").Exists()) + require.Equal(t, beforeTools.Raw, gjson.GetBytes(body, "tools").Raw) + require.Equal(t, beforeChoice.Exists(), gjson.GetBytes(body, "tool_choice").Exists()) + require.Equal(t, beforeChoice.Raw, gjson.GetBytes(body, "tool_choice").Raw) + }) + } +} + +func TestApplyGrokCacheIdentityUsesPreSanitizationToolIntent(t *testing.T) { + tests := []struct { + name string + intentBody string + }{ + { + name: "unsupported tools removed by sanitizer", + intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`, + }, + { + name: "tool choice removed with unsupported tool", + intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // This is the shape apply receives after patchGrokResponsesBody has + // removed unsupported tools and their associated tool_choice. + patchedBody := []byte(`{"model":"grok-4.5","input":"hello"}`) + body, err := applyGrokResponsesCacheIdentity(patchedBody, []byte(tt.intentBody), "isolated-id", true) + + require.NoError(t, err) + require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String()) + require.False(t, gjson.GetBytes(body, "tools").Exists()) + require.False(t, gjson.GetBytes(body, "tool_choice").Exists()) + }) + } +} + +func TestApplyGrokCacheIdentityWithoutFreeTierRoutingOnlyWritesIdentity(t *testing.T) { + sourceBody := []byte(`{"model":"grok-4.5","input":"hello"}`) + body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", false) + + require.NoError(t, err) + require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String()) + require.False(t, gjson.GetBytes(body, "tools").Exists()) + require.False(t, gjson.GetBytes(body, "tool_choice").Exists()) +} + +func TestGrokCompactRequestSkipsCacheIdentityAndNativeTools(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(701) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/compact", nil) + body := []byte(`{"model":"grok","input":"compact this","prompt_cache_key":"raw-client-key"}`) + + identity := resolveGrokCacheIdentity(c, body, "", "grok-4.5") + patched, err := applyGrokResponsesCacheIdentity(body, body, identity, true) + + require.NoError(t, err) + require.Empty(t, identity) + require.False(t, gjson.GetBytes(patched, "prompt_cache_key").Exists()) + require.False(t, gjson.GetBytes(patched, "tools").Exists()) + require.False(t, gjson.GetBytes(patched, "tool_choice").Exists()) +} + +func TestResolveGrokCacheIdentityConcurrentDeterminism(t *testing.T) { + gin.SetMode(gin.TestMode) + const workers = 50 + body := []byte(`{"model":"grok","messages":[{"role":"system","content":"stable"},{"role":"user","content":"hello"}]}`) + identities := make(chan string, workers) + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + identities <- resolveGrokCacheIdentity(newGrokCacheTestContext(501), body, "", "grok-4.5") + }() + } + wg.Wait() + close(identities) + + var first string + for identity := range identities { + if first == "" { + first = identity + } + require.Equal(t, first, identity) + } + require.NotEmpty(t, first) +} diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge.go b/backend/internal/service/openai_gateway_grok_chat_bridge.go new file mode 100644 index 0000000000..fd6b426536 --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_chat_bridge.go @@ -0,0 +1,320 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" +) + +const ( + grokChatResponsesEndpoint = "/v1/responses" + grokChatRawEndpoint = "/v1/chat/completions" +) + +var grokChatResponsesBridgeTopLevelFields = map[string]struct{}{ + "model": {}, + "messages": {}, + "stream": {}, + "stream_options": {}, + "max_tokens": {}, + "max_completion_tokens": {}, + "temperature": {}, + "top_p": {}, + "prompt_cache_key": {}, + "tools": {}, + "tool_choice": {}, + "functions": {}, + "function_call": {}, +} + +// grokChatResponsesBridgeEligibility deliberately accepts only request shapes +// whose Chat Completions semantics are preserved by the Responses bridge. +// Everything else stays on raw Chat Completions rather than being silently +// dropped or rewritten. +func grokChatResponsesBridgeEligibility(body []byte) (bool, string) { + var root map[string]json.RawMessage + if err := json.Unmarshal(body, &root); err != nil || root == nil { + return false, "invalid_json" + } + + for _, field := range []string{"stop", "reasoning_effort"} { + if _, exists := root[field]; exists { + return false, "unsupported_" + field + } + } + for _, field := range []string{"tools", "functions"} { + if raw, exists := root[field]; exists && !grokChatNullOrEmptyArray(raw) { + return false, "unsupported_" + field + } + } + if raw, exists := root["tool_choice"]; exists && !grokChatNullOrNone(raw) { + return false, "unsupported_tool_choice" + } + if raw, exists := root["function_call"]; exists && !grokChatNullOrNone(raw) { + return false, "unsupported_function_call" + } + for field := range root { + if _, supported := grokChatResponsesBridgeTopLevelFields[field]; !supported { + return false, "unknown_field_" + field + } + } + + var model string + if raw, ok := root["model"]; !ok || json.Unmarshal(raw, &model) != nil || strings.TrimSpace(model) == "" { + return false, "invalid_model" + } + + if raw, ok := root["stream"]; ok { + var stream *bool + if json.Unmarshal(raw, &stream) != nil || stream == nil { + return false, "invalid_stream" + } + } + if raw, ok := root["stream_options"]; ok { + var options map[string]json.RawMessage + if json.Unmarshal(raw, &options) != nil || options == nil { + return false, "invalid_stream_options" + } + for field, value := range options { + if field != "include_usage" { + return false, "unknown_stream_option_" + field + } + var includeUsage *bool + if json.Unmarshal(value, &includeUsage) != nil || includeUsage == nil { + return false, "invalid_stream_include_usage" + } + } + } + + for _, field := range []string{"max_tokens", "max_completion_tokens"} { + if raw, ok := root[field]; ok { + var value *int + if json.Unmarshal(raw, &value) != nil || value == nil || *value < 128 { + return false, "unsafe_" + field + } + } + } + if _, hasMaxTokens := root["max_tokens"]; hasMaxTokens { + if _, hasMaxCompletionTokens := root["max_completion_tokens"]; hasMaxCompletionTokens { + return false, "conflicting_max_tokens" + } + } + for _, field := range []string{"temperature", "top_p"} { + if raw, ok := root[field]; ok { + var value *float64 + if json.Unmarshal(raw, &value) != nil || value == nil { + return false, "invalid_" + field + } + } + } + if raw, ok := root["prompt_cache_key"]; ok { + var key string + if json.Unmarshal(raw, &key) != nil { + return false, "invalid_prompt_cache_key" + } + } + + var messages []map[string]json.RawMessage + rawMessages, ok := root["messages"] + if !ok || json.Unmarshal(rawMessages, &messages) != nil || len(messages) == 0 { + return false, "invalid_messages" + } + for _, message := range messages { + for field := range message { + if field != "role" && field != "content" { + return false, "unsafe_message_field_" + field + } + } + var role string + if raw, exists := message["role"]; !exists || json.Unmarshal(raw, &role) != nil { + return false, "invalid_message_role" + } + switch role { + case "system", "user", "assistant": + default: + return false, "unsupported_message_role_" + role + } + var content string + if raw, exists := message["content"]; !exists || json.Unmarshal(raw, &content) != nil { + // Structured content includes image_url and other parts whose exact + // behavior is not guaranteed by this bridge. + return false, "non_text_message_content" + } + if strings.TrimSpace(content) == "" { + return false, "empty_message_content" + } + } + + return true, "" +} + +func grokChatNullOrEmptyArray(raw json.RawMessage) bool { + if strings.TrimSpace(string(raw)) == "null" { + return true + } + var values []json.RawMessage + return json.Unmarshal(raw, &values) == nil && len(values) == 0 +} + +func grokChatNullOrNone(raw json.RawMessage) bool { + if strings.TrimSpace(string(raw)) == "null" { + return true + } + var value string + return json.Unmarshal(raw, &value) == nil && strings.EqualFold(strings.TrimSpace(value), "none") +} + +func grokChatCacheIntentBody(body []byte) ([]byte, error) { + var root map[string]json.RawMessage + if err := json.Unmarshal(body, &root); err != nil { + return nil, err + } + for _, field := range []string{"tools", "tool_choice", "functions", "function_call"} { + delete(root, field) + } + return json.Marshal(root) +} + +func grokChatResponsesRuntimeEligible(upstreamModel, cacheIdentity string) bool { + return strings.TrimSpace(upstreamModel) == "grok-4.5" && strings.TrimSpace(cacheIdentity) != "" +} + +// forwardGrokChatCompletionsViaResponses converts a strictly compatible Chat +// request into xAI Responses format and reuses the established Responses-to- +// Chat response translators. It intentionally does not run the Codex OAuth +// transform because Grok CLI is a separate upstream protocol. +func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + promptCacheKey string, + defaultMappedModel string, +) (*OpenAIForwardResult, error) { + startTime := time.Now() + + var chatReq apicompat.ChatCompletionsRequest + if err := json.Unmarshal(body, &chatReq); err != nil { + return nil, fmt.Errorf("parse grok chat completions request: %w", err) + } + originalModel := chatReq.Model + clientStream := chatReq.Stream + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + cacheIdentity := resolveGrokCacheIdentity(c, body, promptCacheKey, upstreamModel) + if !grokChatResponsesRuntimeEligible(upstreamModel, cacheIdentity) { + return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) + } + + responsesReq, err := apicompat.ChatCompletionsToResponses(&chatReq) + if err != nil { + return nil, fmt.Errorf("convert grok chat completions to responses: %w", err) + } + responsesReq.Model = upstreamModel + responsesReq.Stream = true + // These fields are useful to Codex but are not needed by the Grok CLI + // protocol. Keep the bridge request as close as possible to native Grok. + responsesReq.Include = nil + responsesReq.Store = nil + + responsesBody, err := json.Marshal(responsesReq) + if err != nil { + return nil, fmt.Errorf("marshal grok responses bridge request: %w", err) + } + responsesBody, err = patchGrokResponsesBody(responsesBody, upstreamModel) + if err != nil { + return nil, fmt.Errorf("patch grok responses bridge request: %w", err) + } + intentBody, err := grokChatCacheIntentBody(body) + if err != nil { + return nil, fmt.Errorf("normalize grok responses bridge tool intent: %w", err) + } + responsesBody, err = applyGrokResponsesCacheIdentity(responsesBody, intentBody, cacheIdentity, true) + if err != nil { + return nil, fmt.Errorf("apply grok responses bridge cache identity: %w", err) + } + + updatedBody, policyErr := s.applyOpenAIFastPolicyToBody(ctx, account, upstreamModel, responsesBody) + if policyErr != nil { + var blocked *OpenAIFastBlockedError + if errors.As(policyErr, &blocked) { + MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied) + writeChatCompletionsError(c, http.StatusForbidden, "permission_error", blocked.Message) + } + return nil, policyErr + } + responsesBody = updatedBody + + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, fmt.Errorf("get grok access token: %w", err) + } + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, cacheIdentity) + releaseUpstreamCtx() + if err != nil { + return nil, fmt.Errorf("build grok responses bridge request: %w", err) + } + SetActualOpenAIUpstreamEndpoint(c, grokChatResponsesEndpoint) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= http.StatusBadRequest { + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if upstreamMsg == "" { + upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), + Kind: "failover", + Message: upstreamMsg, + }) + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + if s.shouldFailoverUpstreamError(resp.StatusCode) { + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) + } + + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + + var result *OpenAIForwardResult + if clientStream { + result, err = s.handleChatStreamingResponse(resp, c, account, originalModel, billingModel, upstreamModel, startTime, len(body)) + } else { + result, err = s.handleChatBufferedStreamingResponse(resp, c, account, originalModel, billingModel, upstreamModel, startTime) + } + if result != nil { + result.UpstreamEndpoint = grokChatResponsesEndpoint + result.ResponseHeaders = resp.Header.Clone() + if result.RequestID == "" { + result.RequestID = firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")) + } + result.ReasoningEffort = extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) + } + return result, err +} diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go new file mode 100644 index 0000000000..0843f1ea42 --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go @@ -0,0 +1,358 @@ +//go:build unit + +package service + +import ( + "bytes" + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestGrokChatResponsesBridgeEligibility(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + body string + want bool + reason string + }{ + { + name: "plain text chat", + body: `{"model":"grok","messages":[{"role":"system","content":"concise"},{"role":"user","content":"hi"}],"stream":false}`, + want: true, + }, + { + name: "safe generation options", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true,"stream_options":{"include_usage":true},"max_completion_tokens":256,"temperature":0.2,"top_p":0.9,"prompt_cache_key":"session","tools":[],"functions":null,"tool_choice":"none"}`, + want: true, + }, + { + name: "stop falls back", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"stop":"done"}`, + reason: "unsupported_stop", + }, + { + name: "developer role falls back", + body: `{"model":"grok","messages":[{"role":"developer","content":"rules"},{"role":"user","content":"hi"}]}`, + reason: "unsupported_message_role_developer", + }, + { + name: "image content falls back", + body: `{"model":"grok","messages":[{"role":"user","content":[{"type":"image_url","image_url":{"url":"data:image/png;base64,QQ=="}}]}]}`, + reason: "non_text_message_content", + }, + { + name: "function tools fall back", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"tools":[{"type":"function","function":{"name":"lookup"}}]}`, + reason: "unsupported_tools", + }, + { + name: "automatic tool choice falls back", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"tools":[],"tool_choice":"auto"}`, + reason: "unsupported_tool_choice", + }, + { + name: "reasoning effort falls back because conversion adds summary", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"reasoning_effort":"high"}`, + reason: "unsupported_reasoning_effort", + }, + { + name: "both token limits fall back", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"max_tokens":256,"max_completion_tokens":256}`, + reason: "conflicting_max_tokens", + }, + { + name: "empty message falls back", + body: `{"model":"grok","messages":[{"role":"assistant","content":""},{"role":"user","content":"hi"}]}`, + reason: "empty_message_content", + }, + { + name: "tool history falls back", + body: `{"model":"grok","messages":[{"role":"assistant","content":"","tool_calls":[]}]}`, + reason: "unsafe_message_field_tool_calls", + }, + { + name: "unknown field falls back", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"seed":7}`, + reason: "unknown_field_seed", + }, + { + name: "small max tokens falls back because conversion clamps it", + body: `{"model":"grok","messages":[{"role":"user","content":"hi"}],"max_tokens":32}`, + reason: "unsafe_max_tokens", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, reason := grokChatResponsesBridgeEligibility([]byte(tt.body)) + require.Equal(t, tt.want, got) + require.Equal(t, tt.reason, reason) + }) + } +} + +func TestGrokChatResponsesRuntimeEligibility(t *testing.T) { + t.Parallel() + require.True(t, grokChatResponsesRuntimeEligible("grok-4.5", "isolated-id")) + require.False(t, grokChatResponsesRuntimeEligible("grok-4.3", "isolated-id")) + require.False(t, grokChatResponsesRuntimeEligible("grok-4.5-build-free", "isolated-id")) + require.False(t, grokChatResponsesRuntimeEligible("grok-4.5", "")) +} + +func TestForwardGrokChatViaResponsesNonStreamingCachesAndReturnsChat(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"grok","messages":[{"role":"system","content":"be concise"},{"role":"user","content":"hi"}],"stream":false,"prompt_cache_key":"stable-session","tools":[],"functions":null,"tool_choice":"none"}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 7101}) + + account := grokChatBridgeTestAccount(71) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: grokChatBridgeCompletedResponse("resp_grok_chat_cache", 9856)} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) + require.Equal(t, grokChatResponsesEndpoint, result.UpstreamEndpoint) + require.Equal(t, "grok-4.5", result.UpstreamModel) + require.Equal(t, 9908, result.Usage.InputTokens) + require.Equal(t, 12, result.Usage.OutputTokens) + require.Equal(t, 9856, result.Usage.CacheReadInputTokens) + + identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String() + require.NotEmpty(t, identity) + require.NotEqual(t, "stable-session", identity) + require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) + require.Equal(t, grokFreeCacheDisabledToolChoice, gjson.GetBytes(upstream.lastBody, "tool_choice").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.Equal(t, "system", gjson.GetBytes(upstream.lastBody, "input.0.role").String()) + require.Equal(t, "user", gjson.GetBytes(upstream.lastBody, "input.1.role").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "instructions").Exists()) + require.False(t, gjson.GetBytes(upstream.lastBody, "include").Exists()) + require.False(t, gjson.GetBytes(upstream.lastBody, "store").Exists()) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Equal(t, "cached ok", gjson.Get(recorder.Body.String(), "choices.0.message.content").String()) + require.Equal(t, int64(9856), gjson.Get(recorder.Body.String(), "usage.prompt_tokens_details.cached_tokens").Int()) + require.NotNil(t, repo.updates[account.ID][grokQuotaSnapshotExtraKey]) +} + +func TestForwardGrokChatViaResponsesStreamingPropagatesCachedUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 7201}) + + account := grokChatBridgeTestAccount(72) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: grokChatBridgeCompletedResponse("resp_grok_chat_stream", 4096)} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Stream) + require.Equal(t, grokChatResponsesEndpoint, result.UpstreamEndpoint) + require.Equal(t, 4096, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream") + require.Contains(t, recorder.Body.String(), `"content":"cached ok"`) + require.Contains(t, recorder.Body.String(), `"cached_tokens":4096`) + require.Contains(t, recorder.Body.String(), "data: [DONE]") +} + +func TestForwardGrokChatRuntimeGateFallsBackToRaw(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + setAPIKey bool + mappedModel string + wantUpstream string + }{ + {name: "missing cache identity", wantUpstream: "grok-4.5"}, + {name: "non cache capable mapped model", setAPIKey: true, mappedModel: "grok-4.3", wantUpstream: "grok-4.3"}, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + if tt.setAPIKey { + c.Set("api_key", &APIKey{ID: int64(7301 + index)}) + } + + account := grokChatBridgeTestAccount(int64(73 + index)) + if tt.mappedModel != "" { + account.Credentials["model_mapping"] = map[string]any{"grok": tt.mappedModel} + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"chat_raw","object":"chat.completion","model":"` + tt.wantUpstream + `","choices":[{"index":0,"message":{"role":"assistant","content":"raw ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":2,"completion_tokens":1,"total_tokens":3}}`, + )), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, grokChatRawEndpoint, result.UpstreamEndpoint) + require.Equal(t, tt.wantUpstream, result.UpstreamModel) + require.False(t, gjson.GetBytes(upstream.lastBody, "tools").Exists()) + require.Equal(t, "raw ok", gjson.Get(recorder.Body.String(), "choices.0.message.content").String()) + }) + } +} + +func TestForwardGrokChatViaResponses429UsesGrokRateLimitPolicy(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 7501}) + + account := grokChatBridgeTestAccount(75) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Retry-After": []string{"45"}, + }, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + before := time.Now() + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.Error(t, err) + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.True(t, errors.As(err, &failoverErr)) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) + require.Equal(t, grokChatResponsesEndpoint, GetActualOpenAIUpstreamEndpoint(c)) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Zero(t, repo.tempUnschedCalls) + require.WithinDuration(t, before.Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestForwardGrokRawChatErrorRecordsActualEndpoint(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false,"stop":"done"}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 7601}) + + account := grokChatBridgeTestAccount(76) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"bad request"}}`)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.Error(t, err) + require.Nil(t, result) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, grokChatRawEndpoint, GetActualOpenAIUpstreamEndpoint(c)) +} + +func grokChatBridgeTestAccount(id int64) *Account { + return &Account{ + ID: id, + Name: "grok-cache-bridge", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } +} + +func grokChatBridgeCompletedResponse(responseID string, cachedTokens int) *http.Response { + body := strings.Join([]string{ + `data: {"type":"response.output_text.delta","sequence_number":0,"delta":"cached ok"}`, + "", + `data: {"type":"response.completed","sequence_number":1,"response":{"id":"` + responseID + `","object":"response","model":"grok-4.5","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"cached ok"}]}],"usage":{"input_tokens":9908,"output_tokens":12,"total_tokens":9920,"input_tokens_details":{"cached_tokens":` + strconv.Itoa(cachedTokens) + `}}}}`, + "", + }, "\n") + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "Xai-Request-Id": []string{responseID + "-request"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"9"}, + }, + Body: io.NopCloser(strings.NewReader(body)), + } +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 8dbd9ddad0..3b64cbc16b 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -6,6 +6,7 @@ import ( "bytes" "context" "encoding/json" + "fmt" "io" "mime/multipart" "net/http" @@ -41,6 +42,50 @@ func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testi require.Equal(t, "high", gjson.GetBytes(patched, "reasoning.effort").String()) } +func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + upstreamModel string + wantReasoning bool + }{ + {name: "composer fast", upstreamModel: "grok-composer-2.5-fast"}, + {name: "composer shorthand", upstreamModel: "grok-composer"}, + {name: "composer legacy alias", upstreamModel: "composer-2.5"}, + {name: "provider-prefixed composer", upstreamModel: "xai/grok-composer-2.5-fast"}, + {name: "grok 4.5", upstreamModel: "grok-4.5", wantReasoning: true}, + } + + body := []byte(`{ + "model": "grok", + "input": "hello", + "reasoning": {"effort": "medium", "summary": "auto"}, + "reasoning_effort": "medium", + "reasoningEffort": "medium" + }`) + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + patched, err := patchGrokResponsesBody(body, tt.upstreamModel) + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.Equal(t, tt.upstreamModel, gjson.GetBytes(patched, "model").String()) + + if tt.wantReasoning { + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning.effort").String()) + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning_effort").String()) + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoningEffort").String()) + return + } + + require.False(t, gjson.GetBytes(patched, "reasoning").Exists()) + require.False(t, gjson.GetBytes(patched, "reasoning_effort").Exists()) + require.False(t, gjson.GetBytes(patched, "reasoningEffort").Exists()) + }) + } +} + func TestExtractGrokResponsesReasoningEffortSupportsOpenAICompatibleField(t *testing.T) { t.Parallel() @@ -161,6 +206,45 @@ func TestPatchGrokResponsesBodyDropsToolChoiceWhenNoSupportedToolsRemain(t *test require.False(t, gjson.GetBytes(patched, "tool_choice").Exists()) } +func TestPatchGrokResponsesBodyDropsCodexAdditionalToolsInputItems(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok", + "input": [ + { + "type": "additional_tools", + "role": "developer", + "tools": [ + {"type": "namespace", "name": "image_gen"}, + {"type": "function", "name": "wait"} + ] + }, + { + "type": "message", + "role": "developer", + "content": [{"type": "input_text", "text": "system prompt"}] + }, + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "hello"}] + } + ] + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.5") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.Equal(t, "grok-4.5", gjson.GetBytes(patched, "model").String()) + require.Equal(t, 2, len(gjson.GetBytes(patched, "input").Array())) + require.False(t, gjson.GetBytes(patched, `input.#(type=="additional_tools")`).Exists()) + require.Equal(t, "developer", gjson.GetBytes(patched, "input.0.role").String()) + require.Equal(t, "system prompt", gjson.GetBytes(patched, "input.0.content.0.text").String()) + require.Equal(t, "user", gjson.GetBytes(patched, "input.1.role").String()) + require.Equal(t, "hello", gjson.GetBytes(patched, "input.1.content.0.text").String()) +} + func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") @@ -172,20 +256,37 @@ func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) }, } - req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token") + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "isolated-cache-id") require.NoError(t, err) require.Equal(t, http.MethodPost, req.Method) require.Equal(t, "https://xai.test/v1/responses", req.URL.String()) require.Equal(t, "Bearer access-token", req.Header.Get("Authorization")) require.Equal(t, "application/json", req.Header.Get("Content-Type")) require.Contains(t, req.Header.Get("Accept"), "text/event-stream") + require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version")) + require.Equal(t, "isolated-cache-id", req.Header.Get(grokConversationIDHeader)) data, err := io.ReadAll(req.Body) require.NoError(t, err) require.Equal(t, `{"model":"grok-4.3"}`, strings.TrimSpace(string(data))) } -func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) { +func TestBuildGrokResponsesRequestAllowsPublicAPIKeyBaseURLByDefault(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://grok.example.test/v1/", + }, + } + + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "api-key", "") + require.NoError(t, err) + require.Equal(t, "https://grok.example.test/v1/responses", req.URL.String()) + require.Equal(t, "Bearer api-key", req.Header.Get("Authorization")) +} + +func TestBuildGrokResponsesRequestPinsOAuthCustomBaseURLByDefault(t *testing.T) { t.Parallel() account := &Account{ @@ -196,9 +297,9 @@ func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) { }, } - _, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token") - require.Error(t, err) - require.Contains(t, err.Error(), "invalid base url") + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "") + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", req.URL.String()) } func TestGrokMediaGenerationGateCoversImagesAndVideo(t *testing.T) { @@ -324,6 +425,7 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type")) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.JSONEq(t, `{"model":"grok-imagine-image-quality","prompt":"draw a cat"}`, string(upstream.lastBody)) require.Equal(t, http.StatusOK, recorder.Code) require.JSONEq(t, `{"data":[]}`, recorder.Body.String()) @@ -509,6 +611,42 @@ func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T) require.Equal(t, VideoBillingDefaultDurationSeconds, result.VideoDurationSeconds) } +func TestForwardGrokMediaOAuthImageToVideoUsesOfficialAPIForLargeBody(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + imageData := strings.Repeat("A", 2*1024*1024) + body := []byte(`{"model":"grok-imagine-video-1.5","prompt":"animate","image":{"image_url":"data:image/png;base64,` + imageData + `"}}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/videos/generations", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 66, + Name: "grok-oauth", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-access-token", + "base_url": xai.DefaultCLIBaseURL, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + }, + Body: io.NopCloser(strings.NewReader(`{"request_id":"video-request-oauth"}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + _, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/videos/generations", upstream.lastReq.URL.String()) + require.Equal(t, "data:image/png;base64,"+imageData, gjson.GetBytes(upstream.lastBody, "image.image_url").String()) +} + func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) @@ -550,6 +688,49 @@ func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) { require.Equal(t, "xai-video-req", result.RequestID) } +func TestForwardGrokMediaVideoMutationEndpoints(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + endpoint GrokMediaEndpoint + path string + }{ + {name: "edit", endpoint: GrokMediaEndpointVideosEdits, path: "/videos/edits"}, + {name: "extension", endpoint: GrokMediaEndpointVideosExtensions, path: "/videos/extensions"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-imagine-video","prompt":"continue","video":{"url":"https://example.com/in.mp4"},"duration":6}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1"+tt.path, bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 71, Name: "grok", Platform: PlatformGrok, Type: AccountTypeAPIKey, Concurrency: 1, + Credentials: map[string]any{"api_key": "api-key", "base_url": "https://xai.test/v1"}, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"request_id":"video-mutation-123"}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.ForwardGrokMedia(context.Background(), c, account, tt.endpoint, "", body, "application/json") + require.NoError(t, err) + require.Equal(t, "https://xai.test/v1"+tt.path, upstream.lastReq.URL.String()) + require.Equal(t, http.MethodPost, upstream.lastReq.Method) + require.JSONEq(t, string(body), string(upstream.lastBody)) + require.Equal(t, "video-mutation-123", result.ResponseID) + require.Equal(t, 1, result.VideoCount) + require.Equal(t, 6, result.VideoDurationSeconds) + }) + } +} + func TestBindGrokMediaVideoRequestAccountUsesRequestIDStickyHash(t *testing.T) { ctx := context.Background() groupID := int64(7) @@ -565,7 +746,7 @@ func TestBindGrokMediaVideoRequestAccountUsesRequestIDStickyHash(t *testing.T) { require.Equal(t, int64(63), accountID) } -func TestForwardGrokMediaErrorHonorsCustomErrorCodes(t *testing.T) { +func TestForwardGrokMedia429ReconcilesRateLimitBeforeCustomErrorBypass(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) @@ -585,18 +766,20 @@ func TestForwardGrokMediaErrorHonorsCustomErrorCodes(t *testing.T) { "api_key": "api-key", "base_url": "https://xai.test/v1", "custom_error_codes_enabled": true, - "custom_error_codes": []any{float64(http.StatusTooManyRequests)}, + "custom_error_codes": []any{float64(http.StatusBadRequest)}, }, } upstream := &httpUpstreamRecorder{resp: &http.Response{ - StatusCode: http.StatusBadRequest, + StatusCode: http.StatusTooManyRequests, Header: http.Header{ "Content-Type": []string{"application/json"}, "Xai-Request-Id": []string{"xai-error-req"}, + "Retry-After": []string{"45"}, }, Body: io.NopCloser(strings.NewReader(`{"error":{"message":"do not expose this upstream detail"}}`)), }} - svc := &OpenAIGatewayService{httpUpstream: upstream} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{httpUpstream: upstream, accountRepo: repo} result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesGenerations, "", body, "application/json") require.Error(t, err) @@ -604,15 +787,19 @@ func TestForwardGrokMediaErrorHonorsCustomErrorCodes(t *testing.T) { require.Equal(t, http.StatusInternalServerError, recorder.Code) require.Contains(t, recorder.Body.String(), "Upstream gateway error") require.NotContains(t, recorder.Body.String(), "do not expose") + require.Equal(t, 1, repo.rateLimitedCalls) + require.Zero(t, repo.tempUnschedCalls) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) } -func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *testing.T) { +func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) - body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`) + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false,"stop":"done","prompt_cache_key":"raw-client-cache-key"}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 5101}) account := &Account{ ID: 51, @@ -641,7 +828,7 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te "X-Ratelimit-Limit-Tokens": []string{"1000"}, "X-Ratelimit-Remaining-Tokens": []string{"990"}, }, - Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2}}`)), + Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2,"prompt_tokens_details":{"cached_tokens":1}}}`)), }} svc := &OpenAIGatewayService{ httpUpstream: upstream, @@ -653,24 +840,29 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.NotEqual(t, "raw-client-cache-key", upstream.lastReq.Header.Get(grokConversationIDHeader)) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").Exists()) require.Equal(t, "grok", result.Model) require.Equal(t, "grok-4.5", result.UpstreamModel) require.Equal(t, 1, result.Usage.InputTokens) require.Equal(t, 2, result.Usage.OutputTokens) + require.Equal(t, 1, result.Usage.CacheReadInputTokens) require.NotNil(t, repo.updates[51][grokQuotaSnapshotExtraKey]) require.Equal(t, http.StatusOK, recorder.Code) } -func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) { +func TestForwardGrokResponsesStreamingDefaultsEmptyModelTo45AndSnapshots(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) - body := []byte(`{"model":"grok","input":"hi","stream":true,"reasoning_effort":"high"}`) + body := []byte(`{"input":"hi","stream":true,"reasoning_effort":"high"}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") c.Request.Header.Set("OpenAI-Beta", "responses=experimental") + c.Set("api_key", &APIKey{ID: 5201}) account := &Account{ ID: 52, @@ -713,12 +905,17 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) accountRepo: repo, } - result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "", true, time.Now()) require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()) + require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) + require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String()) require.Equal(t, "high", gjson.GetBytes(upstream.lastBody, "reasoning_effort").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) require.True(t, result.Stream) @@ -734,6 +931,83 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) require.NotNil(t, repo.updates[52][grokQuotaSnapshotExtraKey]) } +func TestForwardGrokResponsesAPIKeyUsesXAIResponses(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","input":"hi","stream":true}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 53, + Name: "grok-api-key", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 2, + Credentials: map[string]any{ + "api_key": "xai-test-key", + "base_url": "https://api.x.ai/v1", + }, + } + upstreamBody := strings.Join([]string{ + `data: {"type":"response.output_text.delta","sequence_number":0,"delta":"ok"}`, + "", + `data: {"type":"response.completed","sequence_number":1,"response":{"id":"resp_grok_api_key","model":"grok-4.5","usage":{"input_tokens":2,"output_tokens":1}}}`, + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer xai-test-key", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "resp_grok_api_key", result.ResponseID) + require.Equal(t, 2, result.Usage.InputTokens) + require.Equal(t, 1, result.Usage.OutputTokens) +} + +func TestAccountTestServiceGrokAPIKeyUsesXAIResponses(t *testing.T) { + gin.SetMode(gin.TestMode) + + account := &Account{ + ID: 54, + Name: "grok-api-key", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 2, + Credentials: map[string]any{ + "api_key": "xai-test-key", + "base_url": "https://api.x.ai/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" + + "data: {\"type\":\"response.completed\"}\n\n", + )), + }} + svc := &AccountTestService{httpUpstream: upstream} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/54/test", nil) + + err := svc.testGrokAccountConnection(c, account, "grok") + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer xai-test-key", upstream.lastReq.Header.Get("Authorization")) + require.Contains(t, recorder.Body.String(), `"type":"test_complete"`) +} + func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) { gin.SetMode(gin.TestMode) @@ -801,6 +1075,205 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey]) } +func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","input":"hi","stream":false,"tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Set("api_key", &APIKey{ID: 5202}) + + account := &Account{ + ID: 56, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{56: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Xai-Request-Id": []string{"xai-non-stream-req"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_grok_non_stream","object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":7,"output_tokens":2,"total_tokens":9,"input_tokens_details":{"cached_tokens":4}}}`)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", false, time.Now()) + require.NoError(t, err) + require.NotNil(t, result) + require.False(t, result.Stream) + require.Equal(t, "resp_grok_non_stream", result.ResponseID) + require.Equal(t, 7, result.Usage.InputTokens) + require.Equal(t, 2, result.Usage.OutputTokens) + require.Equal(t, 4, result.Usage.CacheReadInputTokens) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String() + require.NotEmpty(t, identity) + require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader)) + // The sanitizer drops this unsupported client tool, but its explicit intent + // must still prevent native cache-routing tools from being injected. + require.False(t, gjson.GetBytes(upstream.lastBody, "tools").Exists()) + require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists()) + require.Equal(t, "resp_grok_non_stream", gjson.Get(recorder.Body.String(), "id").String()) +} + +func TestForwardGrokResponsesFailoverKeepsCacheIdentityAcrossAccounts(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","input":[{"role":"user","content":"stable prefix"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 5203}) + + newAccount := func(id int64, token string) *Account { + return &Account{ + ID: id, + Name: fmt.Sprintf("grok-%d", id), + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": token, + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + } + firstAccount := newAccount(58, "access-token-a") + secondAccount := newAccount(59, "access-token-b") + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{58: firstAccount, 59: secondAccount}, + }, + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusServiceUnavailable, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"temporary"}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_after_failover","object":"response","model":"grok-4.3","status":"completed","output":[],"usage":{"input_tokens":5,"output_tokens":1}}`)), + }, + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + _, err := svc.forwardGrokResponses(context.Background(), c, firstAccount, body, "grok", false, time.Now()) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + + result, err := svc.forwardGrokResponses(context.Background(), c, secondAccount, body, "grok", false, time.Now()) + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.requests, 2) + require.Len(t, upstream.bodies, 2) + firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String() + secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String() + require.NotEmpty(t, firstIdentity) + require.Equal(t, firstIdentity, secondIdentity) + require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader)) + require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader)) + require.Equal(t, "Bearer access-token-a", upstream.requests[0].Header.Get("Authorization")) + require.Equal(t, "Bearer access-token-b", upstream.requests[1].Header.Get("Authorization")) +} + +func TestForwardAsChatCompletionsForGrokStreamingStopFallsBackToRawXAIChatCompletions(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true,"stop":"done"}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set(grokConversationIDHeader, "native-client-conversation") + c.Set("api_key", &APIKey{ID: 5301}) + + account := &Account{ + ID: 53, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{53: account}, + }, + } + upstreamBody := strings.Join([]string{ + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[{"index":0,"delta":{"content":"ok"}}]}`, + "", + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":6,"completion_tokens":4,"total_tokens":10,"prompt_tokens_details":{"cached_tokens":1}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"chat-stream-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"7"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) + require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.NotEqual(t, "native-client-conversation", upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool()) + require.True(t, result.Stream) + require.Equal(t, 6, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.OutputTokens) + require.Equal(t, 1, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Body.String(), "data: [DONE]") + require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey]) +} + func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T) { gin.SetMode(gin.TestMode) @@ -809,6 +1282,7 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T) body := []byte(`{"model":"grok-composer-2.5-fast","messages":[{"role":"system","content":"You are concise."},{"role":"user","content":[{"type":"text","text":"What is shown?"},{"type":"image_url","image_url":{"url":"data:image/png;base64,QUJD"}}]}],"stream":false}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") + c.Set("api_key", &APIKey{ID: 5501}) account := &Account{ ID: 55, @@ -858,9 +1332,11 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T) require.NotNil(t, result) require.Len(t, upstream.requests, 2) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.requests[0].URL.String()) + require.Empty(t, upstream.requests[0].Header.Get(grokConversationIDHeader)) require.Equal(t, "grok-build-0.1", gjson.GetBytes(upstream.bodies[0], "model").String()) require.Equal(t, "input_image", gjson.GetBytes(upstream.bodies[0], "input.0.content.1.type").String()) require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.requests[1].URL.String()) + require.NotEmpty(t, upstream.requests[1].Header.Get(grokConversationIDHeader)) require.Equal(t, "grok-composer-2.5-fast", gjson.GetBytes(upstream.bodies[1], "model").String()) require.False(t, strings.Contains(string(upstream.bodies[1]), "image_url")) require.Contains(t, gjson.GetBytes(upstream.bodies[1], "messages.1.content").String(), "Image 1 description") @@ -878,6 +1354,9 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { c, _ := gin.CreateTestContext(recorder) body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 5401}) + c.Request.Header.Set("OpenAI-Beta", "grok-experimental") + c.Request.Header.Set("originator", "opencode") account := &Account{ ID: 54, @@ -896,7 +1375,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { accountsByID: map[int64]*Account{54: account}, }, } - upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_grok_messages", "grok-4.3")} + upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages", 3)} svc := &OpenAIGatewayService{ httpUpstream: upstream, grokTokenProvider: NewGrokTokenProvider(repo, nil), @@ -908,18 +1387,88 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) + require.Empty(t, upstream.lastReq.Header.Get("originator")) + require.Empty(t, upstream.lastReq.Header.Get("version")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()) + require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) + require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String()) + require.Empty(t, upstream.lastReq.Header.Get("session_id")) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) require.NotContains(t, string(upstream.lastBody), "chatgpt.com") require.Equal(t, "grok", result.Model) require.Equal(t, "grok-4.5", result.UpstreamModel) require.Equal(t, 5, result.Usage.InputTokens) require.Equal(t, 2, result.Usage.OutputTokens) + require.Equal(t, 3, result.Usage.CacheReadInputTokens) require.Contains(t, recorder.Body.String(), `"type":"message"`) + require.Equal(t, int64(3), gjson.Get(recorder.Body.String(), "usage.cache_read_input_tokens").Int()) require.Contains(t, recorder.Body.String(), "ok") } -func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) { +func TestForwardAsAnthropicForGrokStreamingPreservesCacheUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","max_tokens":32,"stream":true,"messages":[{"role":"user","content":"hi"}]}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 5402}) + + account := &Account{ + ID: 57, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{57: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages_stream", 2)} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) + identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String() + require.NotEmpty(t, identity) + require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream") + require.Contains(t, recorder.Body.String(), `"cache_read_input_tokens":2`) +} + +func grokMessagesSSECompletedResponse(responseID string, cachedTokens int) *http.Response { + body := strings.Join([]string{ + fmt.Sprintf(`data: {"type":"response.completed","response":{"id":%q,"object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7,"input_tokens_details":{"cached_tokens":%d}}}}`, responseID, cachedTokens), + "", + "data: [DONE]", + "", + }, "\n") + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func TestHandleGrokAccountUpstreamErrorTempUnschedulesNonRateLimitStates(t *testing.T) { tests := []struct { name string status int @@ -931,24 +1480,23 @@ func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing { name: "unauthorized reauth", status: http.StatusUnauthorized, - wantReason: "grok oauth token unauthorized", + wantReason: "grok credentials unauthorized", wantMinCooldown: 10*time.Minute - time.Second, wantMaxCooldown: 10*time.Minute + time.Second, }, { name: "forbidden entitlement", status: http.StatusForbidden, - wantReason: "grok entitlement or subscription tier denied", + wantReason: "grok access or entitlement denied", wantMinCooldown: 30*time.Minute - time.Second, wantMaxCooldown: 30*time.Minute + time.Second, }, { - name: "rate limited retry after", - status: http.StatusTooManyRequests, - headers: http.Header{"Retry-After": []string{"45"}}, - wantReason: "grok rate limited", - wantMinCooldown: 44 * time.Second, - wantMaxCooldown: 46 * time.Second, + name: "upstream temporary error", + status: http.StatusInternalServerError, + wantReason: "grok upstream temporary error", + wantMinCooldown: 2*time.Minute - time.Second, + wantMaxCooldown: 2*time.Minute + time.Second, }, } @@ -963,6 +1511,7 @@ func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) require.Equal(t, 1, repo.tempUnschedCalls) + require.Zero(t, repo.rateLimitedCalls) require.Equal(t, account.ID, repo.lastTempUnschedID) require.Equal(t, tt.wantReason, repo.lastTempUnschedReason) require.True(t, repo.lastTempUnschedUntil.After(before.Add(tt.wantMinCooldown))) @@ -971,10 +1520,83 @@ func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing } } -func TestHandleGrokAccountUpstreamErrorDoesNotShortenExistingPause(t *testing.T) { +func TestHandleGrokAccountUpstreamError429SetsRateLimitedFromRetryAfter(t *testing.T) { + account := &Account{ID: 61, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + before := time.Now() + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, http.Header{"Retry-After": []string{"45"}}, nil) + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, before.Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) +} + +func TestHandleGrokAccountUpstreamError429UsesLatestExhaustedWindowReset(t *testing.T) { + now := time.Now() + requestReset := now.Add(10 * time.Minute).Truncate(time.Second) + tokenReset := now.Add(20 * time.Minute).Truncate(time.Second) + headers := http.Header{ + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"0"}, + "X-Ratelimit-Reset-Requests": []string{fmt.Sprintf("%d", requestReset.Unix())}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"0"}, + "X-Ratelimit-Reset-Tokens": []string{fmt.Sprintf("%d", tokenReset.Unix())}, + } + account := &Account{ID: 62, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, headers, nil) + + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, tokenReset, repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) +} + +func TestHandleGrokAccountUpstreamError429UsesFallbackReset(t *testing.T) { + account := &Account{ID: 63, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + before := time.Now() + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, nil, nil) + + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, before.Add(grokRateLimitFallbackCooldown), repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) +} + +func TestGrokRateLimitResetAtUsesFutureWindowAfterRetryAfterExpires(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + observedAt := now.Add(-2 * time.Minute) + windowReset := now.Add(15 * time.Minute) + retryAfter := 30 + snapshot := &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + UpdatedAt: observedAt.Format(time.RFC3339), + RetryAfterSeconds: &retryAfter, + Requests: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(10), + Remaining: grokInt64PtrForTest(0), + ResetUnix: grokInt64PtrForTest(windowReset.Unix()), + }, + } + + resetAt, limited := grokRateLimitResetAt(snapshot, now) + + require.True(t, limited) + require.WithinDuration(t, windowReset, resetAt, time.Second) +} + +func TestHandleGrokAccountUpstreamError429DoesNotShortenExistingPause(t *testing.T) { existingUntil := time.Now().Add(15 * time.Minute) account := &Account{ - ID: 62, + ID: 64, Platform: PlatformGrok, Type: AccountTypeOAuth, TempUnschedulableUntil: &existingUntil, @@ -985,11 +1607,234 @@ func TestHandleGrokAccountUpstreamErrorDoesNotShortenExistingPause(t *testing.T) svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, http.Header{"Retry-After": []string{"45"}}, nil) - require.Equal(t, 1, repo.tempUnschedCalls) - require.WithinDuration(t, existingUntil, repo.lastTempUnschedUntil, time.Second) + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID) require.True(t, ok) runtimeUntil, ok := value.(time.Time) require.True(t, ok) require.WithinDuration(t, existingUntil, runtimeUntil, time.Second) } + +func TestUpdateGrokUsageSnapshotExhaustedSuccessBypassesThrottleAndSetsRateLimited(t *testing.T) { + account := &Account{ID: 65, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{ + accountRepo: repo, + codexSnapshotThrottle: newAccountWriteThrottle(time.Hour), + } + now := time.Now() + + // Consume the normal snapshot write allowance first. + svc.updateGrokUsageSnapshot(context.Background(), account, &xai.QuotaSnapshot{ + StatusCode: http.StatusOK, + Requests: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(10), + Remaining: grokInt64PtrForTest(9), + }, + UpdatedAt: now.UTC().Format(time.RFC3339), + }) + resetAt := now.Add(30 * time.Minute).Truncate(time.Second) + svc.updateGrokUsageSnapshot(context.Background(), account, &xai.QuotaSnapshot{ + StatusCode: http.StatusOK, + Requests: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(10), + Remaining: grokInt64PtrForTest(0), + ResetUnix: grokInt64PtrForTest(resetAt.Unix()), + ResetAt: resetAt.UTC().Format(time.RFC3339), + }, + UpdatedAt: now.UTC().Format(time.RFC3339), + }) + + require.Equal(t, 2, repo.updateCalls) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, resetAt, repo.lastRateLimitResetAt, time.Second) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestUpdateGrokUsageSnapshotAvailableSuccessDoesNotSetRateLimited(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 66, Platform: PlatformGrok, Type: AccountTypeOAuth} + + svc.updateGrokUsageSnapshot(context.Background(), account, &xai.QuotaSnapshot{ + StatusCode: http.StatusOK, + Requests: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(10), + Remaining: grokInt64PtrForTest(1), + }, + UpdatedAt: time.Now().UTC().Format(time.RFC3339), + }) + + require.Equal(t, 1, repo.updateCalls) + require.Zero(t, repo.rateLimitedCalls) +} + +func TestUpdateGrokUsageSnapshotExhaustedSuccessWithoutResetUsesFallback(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 67, Platform: PlatformGrok, Type: AccountTypeOAuth} + before := time.Now() + + svc.updateGrokUsageSnapshot(context.Background(), account, &xai.QuotaSnapshot{ + StatusCode: http.StatusOK, + Tokens: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(2_000_000), + Remaining: grokInt64PtrForTest(0), + }, + UpdatedAt: before.UTC().Format(time.RFC3339), + }) + + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, before.Add(grokRateLimitFallbackCooldown), repo.lastRateLimitResetAt, time.Second) + stored, ok := repo.updates[account.ID][grokQuotaSnapshotExtraKey].(*xai.QuotaSnapshot) + require.True(t, ok) + require.NotNil(t, stored.Tokens.ResetUnix) + paused, _ := shouldAutoPauseGrokQuotaWindow("tokens", stored.Tokens, before.Add(time.Second)) + require.True(t, paused) + paused, _ = shouldAutoPauseGrokQuotaWindow("tokens", stored.Tokens, repo.lastRateLimitResetAt.Add(time.Second)) + require.False(t, paused) +} + +func TestOpenAIWSHTTPBridgeGrok429PersistsRateLimit(t *testing.T) { + repo := &grokQuotaAccountRepo{} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"45"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + svc := &OpenAIGatewayService{accountRepo: repo, httpUpstream: upstream} + account := &Account{ID: 68, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1} + before := time.Now() + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), nil, account, "token", + []byte(`{"type":"response.create","model":"grok-4.3","input":"hi"}`), + 64, "grok-4.3", "", "", "", "cache-id", 1, + func([]byte) error { return nil }, + ) + + require.Error(t, err) + require.Nil(t, result) + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, before.Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) + require.Zero(t, repo.tempUnschedCalls) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestOpenAIWSHTTPBridgeGrokExhaustedSuccessPersistsRateLimit(t *testing.T) { + repo := &grokQuotaAccountRepo{} + resetAt := time.Now().Add(20 * time.Minute).UTC().Truncate(time.Second) + resp := grokMessagesSSECompletedResponse("resp_ws_limited", 0) + resp.Header.Set("X-Ratelimit-Limit-Requests", "10") + resp.Header.Set("X-Ratelimit-Remaining-Requests", "0") + resp.Header.Set("X-Ratelimit-Reset-Requests", fmt.Sprintf("%d", resetAt.Unix())) + upstream := &httpUpstreamRecorder{resp: resp} + svc := &OpenAIGatewayService{accountRepo: repo, httpUpstream: upstream} + account := &Account{ID: 69, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1} + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), nil, account, "token", + []byte(`{"type":"response.create","model":"grok-4.3","input":"hi"}`), + 64, "grok-4.3", "", "", "", "cache-id", 1, + func([]byte) error { return nil }, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, resetAt, repo.lastRateLimitResetAt, time.Second) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestFailoverOpenAIUpstreamHTTPErrorUsesOnlyGrokRateLimitPolicy(t *testing.T) { + gin.SetMode(gin.TestMode) + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 70, Platform: PlatformGrok, Type: AccountTypeOAuth} + resp := &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"45"}}, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + + failoverErr := svc.failoverOpenAIUpstreamHTTPError( + context.Background(), c, account, resp, + []byte(`{"error":{"message":"rate limited"}}`), "rate limited", "grok-4.3", + ) + + require.NotNil(t, failoverErr) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Zero(t, repo.tempUnschedCalls) +} + +func TestPatchGrokResponsesBody_StripsReasoningContentNull(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok-latest", + "input": [ + {"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}, + {"type":"reasoning","summary":[{"type":"summary_text","text":"thinking..."}],"content":null,"encrypted_content":null}, + {"type":"message","role":"assistant","content":[{"type":"output_text","text":"Hello!"}]} + ] + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.5") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + + input := gjson.GetBytes(patched, "input") + require.True(t, input.IsArray()) + + items := input.Array() + require.Len(t, items, 3) + + reasoning := items[1] + require.Equal(t, "reasoning", reasoning.Get("type").String()) + require.True(t, reasoning.Get("summary").Exists(), "summary should be preserved") + require.False(t, reasoning.Get("content").Exists(), "content: null should be stripped") +} + +func TestPatchGrokResponsesBody_KeepsReasoningContentNonNull(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok-latest", + "input": [ + {"type":"reasoning","summary":[{"type":"summary_text","text":"ok"}],"content":"real content"} + ] + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.5") + require.NoError(t, err) + + reasoning := gjson.GetBytes(patched, "input.0") + require.Equal(t, "real content", reasoning.Get("content").String(), "non-null content must not be stripped") +} + +func TestPatchGrokResponsesBody_MultipleReasoningContentNull(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok-latest", + "input": [ + {"type":"reasoning","summary":[{"type":"summary_text","text":"r1"}],"content":null}, + {"type":"message","role":"user","content":"hi"}, + {"type":"reasoning","summary":[{"type":"summary_text","text":"r2"}],"content":null} + ] + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.5") + require.NoError(t, err) + + items := gjson.GetBytes(patched, "input").Array() + require.Len(t, items, 3) + + require.False(t, items[0].Get("content").Exists()) + require.False(t, items[2].Get("content").Exists()) +} diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 4d8b14cc28..9a0401c898 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -244,12 +244,17 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return nil, policyErr } responsesBody = updatedBody + grokCacheIdentity := "" if account.Platform == PlatformGrok { + grokCacheIdentity = resolveGrokCacheIdentity(c, responsesBody, promptCacheKey, upstreamModel) patchedBody, patchErr := patchGrokResponsesBody(responsesBody, upstreamModel) if patchErr != nil { return nil, patchErr } - responsesBody = patchedBody + responsesBody, patchErr = applyGrokResponsesCacheIdentity(patchedBody, responsesBody, grokCacheIdentity, account.IsGrokOAuth()) + if patchErr != nil { + return nil, fmt.Errorf("apply grok prompt cache identity: %w", patchErr) + } } // 5. Get access token @@ -261,15 +266,14 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // 6. Build upstream request if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok { // Messages 兼容桥即使 body 未带 todo-guard/prompt_cache_key 标记(如映射到非 - // gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支:不带 - // originator、User-Agent 逐字透传,避免身份收口(issue #3901)误改本路径 - // 刻意最小化的请求形态(下方的 Del(OpenAI-Beta/originator) 兜底保持不变)。 + // gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支,以保留 + // 既有 body/session/conversation 行为。身份头在 post-build 阶段统一恢复。 setOpenAICompatMessagesBridgeContext(c, true) } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) var upstreamReq *http.Request if account.Platform == PlatformGrok { - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity) } else { upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false) } @@ -280,7 +284,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // Override session_id with a deterministic UUID derived from the isolated // session key, ensuring different API keys produce different upstream sessions. - if promptCacheKey != "" { + if account.Platform != PlatformGrok && promptCacheKey != "" { isolatedSessionID := generateSessionUUID(isolateOpenAISessionID(apiKeyID, promptCacheKey)) upstreamReq.Header.Set("session_id", isolatedSessionID) if upstreamReq.Header.Get("conversation_id") != "" { @@ -288,12 +292,16 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } } if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok { - // Anthropic Messages compatibility uses the ChatGPT Codex SSE endpoint. - // Match airgate-openai's request shape: the SSE endpoint does not need - // the Responses experimental beta header, and forcing originator can make - // ChatGPT select a different internal continuation path. - upstreamReq.Header.Del("OpenAI-Beta") - upstreamReq.Header.Del("originator") + // buildUpstreamRequest 保留 Messages bridge 的 body/session 兼容行为,并会先 + // 清除身份头。真正发送前恢复完整 Codex 身份,避免 ChatGPT Codex 上游因缺失 + // originator/OpenAI-Beta 返回 404(issue #3901)。 + ensureCodexIdentityHeaders(upstreamReq.Header) + enforceCodexIdentityHeaders(upstreamReq.Header) + logger.L().Debug("openai messages: upstream identity restored", + zap.Int64("account_id", account.ID), + zap.String("upstream_model", upstreamModel), + zap.Bool("compat_identity_restored", true), + ) } if account.Type == AccountTypeOAuth && promptCacheKey != "" && strings.TrimSpace(c.GetHeader("conversation_id")) == "" { upstreamReq.Header.Del("conversation_id") @@ -324,10 +332,9 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return s.ForwardAsAnthropic(markAgentIdentityTaskRecoveryTried(ctx), c, account, body, promptCacheKey, defaultMappedModel) } if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) } - if previousResponseID != "" && (isOpenAICompatPreviousResponseNotFound(resp.StatusCode, upstreamMsg, respBody) || isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody)) { if isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody) { s.disableOpenAICompatSessionContinuation(ctx, c, account, promptCacheKey) @@ -347,6 +354,9 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // Non-failover error: return Anthropic-formatted error to client return s.handleAnthropicErrorResponse(resp, c, account, billingModel) } + if account.Platform == PlatformGrok && account.Type == AccountTypeOAuth && !account.IsShadow() { + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + } if account.Type == AccountTypeOAuth && promptCacheKey != "" { if turnState := strings.TrimSpace(resp.Header.Get("x-codex-turn-state")); turnState != "" { @@ -394,10 +404,8 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // Extract and save Codex usage snapshot from response headers (for OAuth accounts). // 排除 spark 影子:其 codex_* 仅由 QueryUsage(/wham/usage bengalfox)更新(外审第7轮 P1)。 - if handleErr == nil && account.Type == AccountTypeOAuth && !account.IsShadow() { - if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) - } else if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { + if handleErr == nil && account.Type == AccountTypeOAuth && !account.IsShadow() && account.Platform != PlatformGrok { + if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { s.updateCodexUsageSnapshot(ctx, account.ID, snapshot) } } diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback.go b/backend/internal/service/openai_gateway_messages_chat_fallback.go index f66fbc9ada..a1a7023be4 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback.go @@ -101,7 +101,7 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions( if err != nil { return nil, err } - resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent()) + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "") if err != nil { return nil, err } @@ -140,7 +140,7 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsAnthropic( if err != nil { return nil, err } - responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel) + responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel, nil, false, nil) anthropicResp := apicompat.ResponsesToAnthropic(responsesResp, originalModel) diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback_test.go b/backend/internal/service/openai_gateway_messages_chat_fallback_test.go index 8bb15c81aa..b36b76ffa8 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback_test.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback_test.go @@ -356,6 +356,8 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint( c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", "third-party-client/1.0.0") + c.Request.Header.Set("originator", "opencode") upstreamBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_native","object":"response","model":"gpt-5.4","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, @@ -385,5 +387,9 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint( "responses-capable account must stay on /v1/responses, got %s", upstream.lastReq.URL.String()) require.True(t, gjson.GetBytes(upstream.lastBody, "input").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "messages").Exists()) + require.Equal(t, "third-party-client/1.0.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "opencode", upstream.lastReq.Header.Get("originator")) + require.Empty(t, upstream.lastReq.Header.Get("version")) + require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) require.Equal(t, "ok", gjson.Get(rec.Body.String(), "content.0.text").String()) } diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index f280f37841..0c0bacea3f 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -943,6 +943,22 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( trimmedData = strings.TrimSpace(string(normalizedData)) line = "data: " + string(normalizedData) } + if normalizedData, normalized := normalizeCompletedImageGenerationStatus(dataBytes); normalized { + dataBytes = normalizedData + trimmedData = strings.TrimSpace(string(normalizedData)) + line = "data: " + string(normalizedData) + } + if trimmedData != "[DONE]" { + restoredData, restoreErr := restoreOpenAIResponsesNamespacePayload(c, dataBytes) + if restoreErr != nil { + return resultWithUsage(), fmt.Errorf("restore OpenAI passthrough namespace response: %w", restoreErr) + } + if !bytes.Equal(restoredData, dataBytes) { + dataBytes = restoredData + trimmedData = strings.TrimSpace(string(restoredData)) + line = "data: " + string(restoredData) + } + } eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String()) if eventType == "response.failed" { failedMessage = extractOpenAISSEErrorMessage(dataBytes) @@ -1123,6 +1139,10 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( if originalModel != "" && mappedModel != "" && originalModel != mappedModel { body = s.replaceModelInResponseBody(body, mappedModel, originalModel) } + body, err = restoreOpenAIResponsesNamespacePayload(c, body) + if err != nil { + return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", err) + } if !writeOpenAICompactSSEBridge(c, resp.StatusCode, body) { c.Data(resp.StatusCode, contentType, body) } @@ -1164,6 +1184,11 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c } // Correct tool calls in final response body = s.correctToolCallsInResponseBody(body) + restoredBody, restoreErr := restoreOpenAIResponsesNamespacePayload(c, body) + if restoreErr != nil { + return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", restoreErr) + } + body = restoredBody } else { terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) if terminalOK && terminalType == "response.failed" { diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 81578f630f..d2eca05d74 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -39,6 +39,17 @@ type openAIRecordUsageBillingRepoStub struct { lastCtxErr error } +type openAIRecordUsageAccountRepoStub struct { + AccountRepository + account *Account + calls int +} + +func (s *openAIRecordUsageAccountRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) { + s.calls++ + return s.account, nil +} + func (s *openAIRecordUsageBillingRepoStub) Apply(ctx context.Context, cmd *UsageBillingCommand) (*UsageBillingApplyResult, error) { s.calls++ s.lastCmd = cmd @@ -1045,7 +1056,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12) } -func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *testing.T) { +func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefault(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} subRepo := &openAIRecordUsageSubRepoStub{} @@ -1063,7 +1074,45 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te }, APIKey: &APIKey{ID: 1014}, User: &User{ID: 2014}, - Account: &Account{ID: 3014}, + Account: &Account{ID: 3014, Platform: PlatformOpenAI}, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + + expectedInput := 300000 * 2.5e-6 + expectedOutput := 2000 * 15e-6 + require.InDelta(t, expectedInput, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10) + require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10) + require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10) + require.False(t, usageRepo.lastLog.LongContextBillingApplied) + require.Equal(t, 1, userRepo.deductCalls) +} + +func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccount(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_gpt54_long_context_disabled", + Usage: OpenAIUsage{ + InputTokens: 300000, + OutputTokens: 2000, + }, + Model: "gpt-5.4-2026-03-05", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 1015}, + User: &User{ID: 2015}, + Account: &Account{ + ID: 3015, + Platform: PlatformOpenAI, + Extra: map[string]any{"openai_long_context_billing_enabled": true}, + }, }) require.NoError(t, err) @@ -1075,7 +1124,62 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10) require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10) - require.Equal(t, 1, userRepo.deductCalls) + require.True(t, usageRepo.lastLog.LongContextBillingApplied) +} + +func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) { + tests := []struct { + name string + parentEnabled bool + }{ + {name: "parent opt out overrides stale enabled shadow", parentEnabled: false}, + {name: "parent opt in overrides stale disabled shadow", parentEnabled: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + accountRepo := &openAIRecordUsageAccountRepoStub{account: &Account{ + ID: 4016, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{openAILongContextBillingEnabledKey: tt.parentEnabled}, + }} + svc := newOpenAIRecordUsageServiceForTest( + usageRepo, + &openAIRecordUsageUserRepoStub{}, + &openAIRecordUsageSubRepoStub{}, + nil, + ) + svc.accountRepo = accountRepo + parentID := int64(4016) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_gpt54_shadow_parent_setting", + Usage: OpenAIUsage{InputTokens: 300000, OutputTokens: 2000}, + Model: "gpt-5.4-2026-03-05", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 1016}, + User: &User{ID: 2016}, + Account: &Account{ + ID: 3016, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + ParentAccountID: &parentID, + QuotaDimension: QuotaDimensionSpark, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: !tt.parentEnabled, + }, + }, + }) + + require.NoError(t, err) + require.Equal(t, 1, accountRepo.calls) + require.Equal(t, tt.parentEnabled, usageRepo.lastLog.LongContextBillingApplied) + }) + } } func TestOpenAIGatewayServiceRecordUsage_ServiceTierPriorityUsesFastPricing(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 935a32f58b..b47c2c8638 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -136,12 +136,20 @@ func sanitizeEncryptedReasoningInputItem(item any) (next any, changed bool, keep return item, false, true } - _, hasEncryptedContent := inputItem["encrypted_content"] - if !hasEncryptedContent { - return item, false, true + if _, has := inputItem["encrypted_content"]; has { + delete(inputItem, "encrypted_content") + changed = true } - delete(inputItem, "encrypted_content") + // xAI 422: "content": null 导致 untagged enum 反序列化失败 + if v, has := inputItem["content"]; has && v == nil { + delete(inputItem, "content") + changed = true + } + + if !changed { + return item, false, true + } if len(inputItem) == 1 { return nil, true, false } @@ -365,15 +373,57 @@ func newOpenAIRequestView(body []byte) openAIRequestView { if len(body) == 0 { return openAIRequestView{} } - return openAIRequestView{ - body: body, - Model: strings.TrimSpace(gjson.GetBytes(body, "model").String()), - Stream: gjson.GetBytes(body, "stream").Bool(), - PromptCacheKey: strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()), - PreviousResponseID: strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()), - ServiceTier: strings.TrimSpace(gjson.GetBytes(body, "service_tier").String()), - ReasoningEffort: strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()), - } + + const ( + modelField uint8 = 1 << iota + streamField + promptCacheKeyField + previousResponseIDField + serviceTierField + reasoningField + allRequestViewFields = modelField | streamField | promptCacheKeyField | + previousResponseIDField | serviceTierField | reasoningField + ) + + view := openAIRequestView{body: body} + var seen uint8 + // parseRawJSONView reads body without copying; view keeps body alive for extracted strings. + parseRawJSONView(body).ForEach(func(key, value gjson.Result) bool { + switch key.Str { + case "model": + if seen&modelField == 0 { + view.Model = strings.TrimSpace(value.String()) + seen |= modelField + } + case "stream": + if seen&streamField == 0 { + view.Stream = value.Bool() + seen |= streamField + } + case "prompt_cache_key": + if seen&promptCacheKeyField == 0 { + view.PromptCacheKey = strings.TrimSpace(value.String()) + seen |= promptCacheKeyField + } + case "previous_response_id": + if seen&previousResponseIDField == 0 { + view.PreviousResponseID = strings.TrimSpace(value.String()) + seen |= previousResponseIDField + } + case "service_tier": + if seen&serviceTierField == 0 { + view.ServiceTier = strings.TrimSpace(value.String()) + seen |= serviceTierField + } + case "reasoning": + if seen&reasoningField == 0 { + view.ReasoningEffort = strings.TrimSpace(value.Get("effort").String()) + seen |= reasoningField + } + } + return seen != allRequestViewFields + }) + return view } // Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。 diff --git a/backend/internal/service/openai_gateway_request_body_reasoning_test.go b/backend/internal/service/openai_gateway_request_body_reasoning_test.go new file mode 100644 index 0000000000..62a8855114 --- /dev/null +++ b/backend/internal/service/openai_gateway_request_body_reasoning_test.go @@ -0,0 +1,112 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestTrimOpenAIEncryptedReasoningItems_ContentNull(t *testing.T) { + reqBody := map[string]any{ + "model": "grok-4.5", + "input": []any{ + map[string]any{"type": "message", "role": "user", "content": "hi"}, + map[string]any{ + "type": "reasoning", + "summary": []any{map[string]any{"type": "summary_text", "text": "thinking..."}}, + "content": nil, + "encrypted_content": nil, + }, + map[string]any{"type": "message", "role": "assistant", "content": "Hello!"}, + }, + } + + changed := trimOpenAIEncryptedReasoningItems(reqBody) + require.True(t, changed) + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + require.Len(t, input, 3) + + reasoning, ok := input[1].(map[string]any) + require.True(t, ok) + assert.Equal(t, "reasoning", reasoning["type"]) + assert.NotNil(t, reasoning["summary"]) + _, hasContent := reasoning["content"] + assert.False(t, hasContent, "content: null should be stripped") + _, hasEncrypted := reasoning["encrypted_content"] + assert.False(t, hasEncrypted, "encrypted_content should be stripped") +} + +func TestTrimOpenAIEncryptedReasoningItems_ContentNullOnly(t *testing.T) { + reqBody := map[string]any{ + "model": "grok-4.5", + "input": []any{ + map[string]any{ + "type": "reasoning", + "summary": []any{map[string]any{"type": "summary_text", "text": "ok"}}, + "content": nil, + }, + }, + } + + changed := trimOpenAIEncryptedReasoningItems(reqBody) + require.True(t, changed) + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + require.Len(t, input, 1) + + reasoning, ok := input[0].(map[string]any) + require.True(t, ok) + _, hasContent := reasoning["content"] + assert.False(t, hasContent, "content: null should be stripped even without encrypted_content") +} + +func TestTrimOpenAIEncryptedReasoningItems_ContentNonNull(t *testing.T) { + reqBody := map[string]any{ + "model": "grok-4.5", + "input": []any{ + map[string]any{ + "type": "reasoning", + "summary": []any{map[string]any{"type": "summary_text", "text": "ok"}}, + "content": "some actual content", + }, + }, + } + + changed := trimOpenAIEncryptedReasoningItems(reqBody) + assert.False(t, changed, "non-null content should not be stripped") + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + reasoning, ok := input[0].(map[string]any) + require.True(t, ok) + assert.Equal(t, "some actual content", reasoning["content"]) +} + +func TestTrimOpenAIEncryptedReasoningItems_NoReasoningItems(t *testing.T) { + reqBody := map[string]any{ + "model": "grok-4.5", + "input": []any{ + map[string]any{"type": "message", "role": "user", "content": "hi"}, + }, + } + + changed := trimOpenAIEncryptedReasoningItems(reqBody) + assert.False(t, changed) +} + +func TestTrimOpenAIEncryptedReasoningItems_ContentNullDropsBareSkeleton(t *testing.T) { + reqBody := map[string]any{ + "input": []any{ + map[string]any{"type": "reasoning", "content": nil}, + }, + } + + changed := trimOpenAIEncryptedReasoningItems(reqBody) + require.True(t, changed) + _, hasInput := reqBody["input"] + assert.False(t, hasInput, "bare reasoning skeleton should be dropped, emptying input") +} diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 5841ce6556..1d70245eb0 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -281,6 +281,11 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp forceFlushFailedEvent = true sawFailedEvent = true } + if normalizedData, normalized := normalizeCompletedImageGenerationStatus(dataBytes); normalized { + dataBytes = normalizedData + data = string(normalizedData) + line = "data: " + data + } imageCounter.AddSSEData(dataBytes) // Correct Codex tool calls if needed (apply_patch -> edit, etc.) @@ -305,6 +310,17 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp line = "data: " + data eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) } + restoredData, restoreErr := restoreOpenAIResponsesNamespacePayload(c, dataBytes) + if restoreErr != nil { + streamEarlyErr = fmt.Errorf("restore OpenAI namespace response: %w", restoreErr) + return + } + if !bytes.Equal(restoredData, dataBytes) { + dataBytes = restoredData + data = string(restoredData) + line = "data: " + data + eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) + } if sanitizedData, sanitized := sanitizeOpenAIResponseFailedEventForClient( dataBytes, eventType, @@ -850,7 +866,10 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r if originalModel != mappedModel { body = s.replaceModelInResponseBody(body, mappedModel, originalModel) } - + body, err = restoreOpenAIResponsesNamespacePayload(c, body) + if err != nil { + return nil, fmt.Errorf("restore OpenAI namespace response: %w", err) + } responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) contentType := "application/json" @@ -920,6 +939,11 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte } // Correct tool calls in final response body = s.correctToolCallsInResponseBody(body) + restoredBody, restoreErr := restoreOpenAIResponsesNamespacePayload(c, body) + if restoreErr != nil { + return nil, fmt.Errorf("restore OpenAI namespace response: %w", restoreErr) + } + body = restoredBody } else { terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) if terminalOK && terminalType == "response.failed" { @@ -1071,6 +1095,9 @@ func extractCodexFinalResponse(body string) ([]byte, bool) { if finalResponse != nil { return } + if normalized, changed := normalizeCompletedImageGenerationStatus(data); changed { + data = normalized + } eventType := gjson.GetBytes(data, "type").String() if eventType == "response.done" || eventType == "response.completed" { if response := gjson.GetBytes(data, "response"); response.Exists() && response.Type == gjson.JSON && response.Raw != "" { @@ -1084,6 +1111,59 @@ func extractCodexFinalResponse(body string) ([]byte, bool) { return nil, false } +func normalizeCompletedImageGenerationStatus(data []byte) ([]byte, bool) { + if len(data) == 0 || !gjson.ValidBytes(data) { + return data, false + } + + shouldNormalize := func(item gjson.Result) bool { + if !item.Exists() || !item.IsObject() || + strings.TrimSpace(item.Get("type").String()) != "image_generation_call" { + return false + } + switch strings.TrimSpace(item.Get("status").String()) { + case "generating", "in_progress": + return strings.TrimSpace(item.Get("result").String()) != "" + default: + return false + } + } + + eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String()) + switch eventType { + case "response.output_item.done": + if !shouldNormalize(gjson.GetBytes(data, "item")) { + return data, false + } + updated, err := sjson.SetBytes(data, "item.status", "completed") + if err != nil { + return data, false + } + return updated, true + case "response.completed", "response.done": + output := gjson.GetBytes(data, "response.output") + if !output.Exists() || !output.IsArray() { + return data, false + } + updated := data + changed := false + for i, item := range output.Array() { + if !shouldNormalize(item) { + continue + } + next, err := sjson.SetBytes(updated, "response.output."+strconv.Itoa(i)+".status", "completed") + if err != nil { + return data, false + } + updated = next + changed = true + } + return updated, changed + default: + return data, false + } +} + func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) { eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String()) switch eventType { @@ -1124,7 +1204,8 @@ func responsesStreamEventMayContributeToOutput(eventType string) bool { } // collectRawResponsesOutputItemsFromSSE 按到达顺序收集 SSE 流中 -// response.output_item.done 携带的原始 item。item 以 raw JSON 逐字节保留, +// response.output_item.done 携带的原始 item。除已产生结果但仍停留在进行中 +// 的图片状态外,item 以 raw JSON 逐字节保留, // 避免经窄结构体重建时丢弃 encrypted_content/summary/opaque 等 compact // 专属或未来新增字段(#3777 问题 2)。若整条流没有任何 done 事件,退回 // 收集 output_item.added 中的 compaction 类 item——compaction 结果没有 @@ -1151,6 +1232,9 @@ func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) { items = append(items, json.RawMessage(item.Raw)) } forEachOpenAISSEDataPayload(bodyText, func(data []byte) { + if normalized, changed := normalizeCompletedImageGenerationStatus(data); changed { + data = normalized + } if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.done" { return } diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index 247c195f24..1082b020dd 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -39,6 +39,18 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( clientStream := responsesReq.Stream serviceTier := extractOpenAIServiceTierFromBody(body) + // custom 工具(如 codex 的 exec)降级为 function 工具转发,回程需按名字还原为 + // custom_tool_call 项,先记下名字集合;tool_search 工具同理,回程还原为 + // tool_search_call 项;namespace 子工具(如 MCP 工具)摊平转发,回程按映射还原 + // 为带 namespace 字段的 function_call 项。 + effectiveTools, err := apicompat.EffectiveResponsesTools(&responsesReq) + if err != nil { + writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return nil, fmt.Errorf("resolve responses tools: %w", err) + } + customTools := apicompat.CustomToolNames(effectiveTools) + toolSearch := apicompat.HasToolSearchTool(effectiveTools) + namespaceTools := apicompat.NamespaceToolNames(effectiveTools) chatReq, err := apicompat.ResponsesToChatCompletionsRequest(&responsesReq) if err != nil { @@ -85,7 +97,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( if err != nil { return nil, err } - resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent()) + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "") if err != nil { return nil, err } @@ -100,15 +112,18 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( } if clientStream { - return s.streamChatCompletionsAsResponses(c, resp, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) + return s.streamChatCompletionsAsResponses(c, resp, originalModel, customTools, toolSearch, namespaceTools, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) } - return s.bufferChatCompletionsAsResponses(c, resp, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) + return s.bufferChatCompletionsAsResponses(c, resp, originalModel, customTools, toolSearch, namespaceTools, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) } func (s *OpenAIGatewayService) bufferChatCompletionsAsResponses( c *gin.Context, resp *http.Response, originalModel string, + customTools map[string]bool, + toolSearch bool, + namespaceTools map[string]apicompat.NamespacedToolName, billingModel string, upstreamModel string, reasoningEffort *string, @@ -120,7 +135,7 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsResponses( if err != nil { return nil, err } - responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel) + responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel, customTools, toolSearch, namespaceTools) if s.responseHeaderFilter != nil { responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) @@ -144,6 +159,9 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( c *gin.Context, resp *http.Response, originalModel string, + customTools map[string]bool, + toolSearch bool, + namespaceTools map[string]apicompat.NamespacedToolName, billingModel string, upstreamModel string, reasoningEffort *string, @@ -154,6 +172,9 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) state := apicompat.NewChatCompletionsToResponsesStreamState(originalModel) + state.CustomTools = customTools + state.ToolSearchDeclared = toolSearch + state.NamespaceTools = namespaceTools clientDisconnected := false writeEvents := func(events []apicompat.ResponsesStreamEvent) { diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 6ac0de47b4..f318adc604 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -23,17 +23,7 @@ import ( // ExtractSessionID extracts the raw session ID from headers or body without hashing. // Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache. func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string { - if c == nil { - return "" - } - sessionID := strings.TrimSpace(c.GetHeader("session_id")) - if sessionID == "" { - sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) - } - if sessionID == "" && len(body) > 0 { - sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) - } - return sessionID + return explicitOpenAIRequestSessionID(c, body) } func explicitOpenAISessionID(c *gin.Context, body []byte) string { @@ -51,11 +41,33 @@ func explicitOpenAISessionID(c *gin.Context, body []byte) string { return sessionID } +// explicitOpenAIRequestSessionID extends the common OpenAI session signals +// with Grok's native conversation header only for requests authenticated to a +// Grok group. This keeps an unrelated x-grok-conv-id header from changing +// scheduling or upstream session behavior for non-Grok groups. +func explicitOpenAIRequestSessionID(c *gin.Context, body []byte) string { + if c == nil { + return "" + } + + sessionID := strings.TrimSpace(c.GetHeader("session_id")) + if sessionID == "" { + sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) + } + if sessionID == "" && isGrokRequestContext(c) { + sessionID = strings.TrimSpace(c.GetHeader(grokConversationIDHeader)) + } + if sessionID == "" && len(body) > 0 { + sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) + } + return sessionID +} + // GenerateExplicitSessionHash generates a sticky-session hash only from explicit // client session signals. It intentionally skips content-derived fallback and is // used by stateless endpoints such as /v1/images. func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body []byte) string { - sessionID := explicitOpenAISessionID(c, body) + sessionID := explicitOpenAIRequestSessionID(c, body) if sessionID == "" { return "" } @@ -70,14 +82,15 @@ func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body // Priority: // 1. Header: session_id // 2. Header: conversation_id -// 3. Body: prompt_cache_key (opencode) -// 4. Body: content-based fallback (model + system + tools + first user message) +// 3. Header: x-grok-conv-id (Grok groups only) +// 4. Body: prompt_cache_key (opencode) +// 5. Body: content-based fallback (model + system + tools + first user message) func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string { if c == nil { return "" } - sessionID := explicitOpenAISessionID(c, body) + sessionID := explicitOpenAIRequestSessionID(c, body) if sessionID == "" && len(body) > 0 { sessionID = deriveOpenAIContentSessionSeed(body) } diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 24f4786e09..dcda4eeab2 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -49,6 +49,7 @@ const ( openAIWSRetryBackoffMaxDefault = 2 * time.Second openAIWSRetryJitterRatioDefault = 0.2 openAICompactSessionSeedKey = "openai_compact_session_seed" + openAIUpstreamEndpointContextKey = "openai_actual_upstream_endpoint" codexCLIVersion = "0.144.1" // Codex 限额快照仅用于后台展示/诊断,不需要每个成功请求都立即落库。 openAICodexSnapshotPersistMinInterval = 30 * time.Second @@ -66,8 +67,10 @@ var openaiAllowedHeaders = map[string]bool{ "user-agent": true, "originator": true, "session_id": true, + "x-codex-beta-features": true, "x-codex-turn-state": true, "x-codex-turn-metadata": true, + responsesLiteHeaderKey: true, } // OpenAI passthrough allowed headers whitelist. @@ -81,8 +84,10 @@ var openaiPassthroughAllowedHeaders = map[string]bool{ "user-agent": true, "originator": true, "session_id": true, + "x-codex-beta-features": true, "x-codex-turn-state": true, "x-codex-turn-metadata": true, + responsesLiteHeaderKey: true, } // codex_cli_only 拒绝时记录的请求头白名单(仅用于诊断日志,不参与上游透传) @@ -223,6 +228,9 @@ type OpenAIForwardResult struct { // UpstreamModel is the actual model sent to the upstream provider after mapping. // Empty when no mapping was applied (requested model was used as-is). UpstreamModel string + // UpstreamEndpoint is the actual upstream API path used for this request. + // It avoids guessing when one downstream protocol can use multiple upstream endpoints. + UpstreamEndpoint string // ServiceTier records the OpenAI Responses API service tier, e.g. "priority" / "flex". // Nil means the request did not specify a recognized tier. ServiceTier *string @@ -246,11 +254,40 @@ type OpenAIForwardResult struct { VideoResolution string // VideoDurationSeconds 是提交时请求的生成时长(xAI 按输出秒数计费),已归一化到 1-15 秒。 VideoDurationSeconds int + // WebSearchCalls 是 Codex alpha/search 网页搜索调用次数(每次成功请求为 1)。 + // 上游不返回 usage 字段,>0 时走按次计费(分组单价 × 次数 × 倍率)。 + WebSearchCalls int wsReplayInput []json.RawMessage wsReplayInputExists bool } +// SetActualOpenAIUpstreamEndpoint records the endpoint selected by the current +// forwarding attempt. It covers error paths where no OpenAIForwardResult is +// available for usage and operations logging. +func SetActualOpenAIUpstreamEndpoint(c *gin.Context, endpoint string) { + if c == nil { + return + } + if endpoint = strings.TrimSpace(endpoint); endpoint != "" { + c.Set(openAIUpstreamEndpointContextKey, endpoint) + } +} + +// GetActualOpenAIUpstreamEndpoint returns the endpoint recorded by the latest +// forwarding attempt in this request. +func GetActualOpenAIUpstreamEndpoint(c *gin.Context) string { + if c == nil { + return "" + } + value, exists := c.Get(openAIUpstreamEndpointContextKey) + if !exists { + return "" + } + endpoint, _ := value.(string) + return strings.TrimSpace(endpoint) +} + type OpenAIWSRetryMetricsSnapshot struct { RetryAttemptsTotal int64 `json:"retry_attempts_total"` RetryBackoffMsTotal int64 `json:"retry_backoff_ms_total"` @@ -370,6 +407,7 @@ type OpenAIGatewayService struct { openaiWSRetryMetrics openAIWSRetryMetrics responseHeaderFilter *responseheaders.CompiledHeaderFilter codexSnapshotThrottle *accountWriteThrottle + codexModelsManifestCache codexModelsManifestCache openaiCompatSessionResponses sync.Map openaiCompatAnthropicDigestSessions sync.Map } diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index 1dde60c9f0..326fde534d 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -27,6 +27,33 @@ func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) { require.Equal(t, "medium", view.ReasoningEffort) } +func TestOpenAIRequestView_ExtractsFieldsAfterLargeInput(t *testing.T) { + body := []byte(`{"model":"gpt-5","input":[{"content":"` + strings.Repeat("payload", 1024) + `"}],"stream":true,"prompt_cache_key":"session-1","previous_response_id":"resp-1","service_tier":"flex","reasoning":{"effort":"high"}}`) + + view := newOpenAIRequestView(body) + + require.Equal(t, "gpt-5", view.Model) + require.True(t, view.Stream) + require.Equal(t, "session-1", view.PromptCacheKey) + require.Equal(t, "resp-1", view.PreviousResponseID) + require.Equal(t, "flex", view.ServiceTier) + require.Equal(t, "high", view.ReasoningEffort) +} + +func TestOpenAIRequestView_KeepsFirstDuplicateField(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5","model":"gpt-5.1","reasoning":{"effort":"low"},"reasoning":{"effort":"high"}}`)) + + require.Equal(t, "gpt-5", view.Model) + require.Equal(t, "low", view.ReasoningEffort) +} + +func TestOpenAIRequestView_KeepsLenientPrefixExtraction(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[`)) + + require.Equal(t, "gpt-5", view.Model) + require.True(t, view.Stream) +} + func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) { view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[{"type":"message","content":"hi"}]}`)) diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index bc14350394..507fb9d5c3 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2961,7 +2961,7 @@ func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) Header: http.Header{"Content-Type": []string{"text/event-stream"}}, } body := []byte(strings.Join([]string{ - `data: {"type":"response.output_item.done","item":{"id":"ig_123","type":"image_generation_call","result":"aGVsbG8=","revised_prompt":"draw a cat","output_format":"png"}}`, + `data: {"type":"response.output_item.done","item":{"id":"ig_123","type":"image_generation_call","status":"generating","result":"aGVsbG8=","revised_prompt":"draw a cat","output_format":"png"}}`, `data: {"type":"response.completed","response":{"id":"resp_img","model":"gpt-5.4","output":[],"usage":{"input_tokens":7,"output_tokens":9,"output_tokens_details":{"image_tokens":4}}}}`, `data: [DONE]`, }, "\n")) @@ -2972,6 +2972,7 @@ func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) require.Equal(t, 4, usage.ImageOutputTokens) require.NotContains(t, rec.Body.String(), "data:") require.Equal(t, "image_generation_call", gjson.Get(rec.Body.String(), "output.0.type").String()) + require.Equal(t, "completed", gjson.Get(rec.Body.String(), "output.0.status").String()) require.Equal(t, "aGVsbG8=", gjson.Get(rec.Body.String(), "output.0.result").String()) require.Equal(t, "draw a cat", gjson.Get(rec.Body.String(), "output.0.revised_prompt").String()) } diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index f96b679cf6..410b4d944d 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -178,7 +178,27 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec if result.ServiceTier != nil { serviceTier = strings.TrimSpace(*result.ServiceTier) } - cost, err = s.calculateOpenAIRecordUsageCost(ctx, result, apiKey, billingModels, multiplier, imageMultiplier, videoMultiplier, tokens, serviceTier) + billingAccount := account + if account.IsShadow() { + billingAccount, err = resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return err + } + } + longContextBillingEnabled := billingAccount.IsOpenAILongContextBillingEnabled() + cost, err = s.calculateOpenAIRecordUsageCost( + ctx, + result, + apiKey, + billingModels, + multiplier, + imageMultiplier, + videoMultiplier, + baseMultiplier, + tokens, + serviceTier, + longContextBillingEnabled, + ) if err != nil { if !isUsagePricingUnavailableError(err) { return err @@ -257,6 +277,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec usageLog.CacheReadCost = cost.CacheReadCost usageLog.TotalCost = cost.TotalCost usageLog.ActualCost = cost.ActualCost + usageLog.LongContextBillingApplied = cost.LongContextBillingApplied } if isVideoUsage && (cost == nil || cost.BillingMode != string(BillingModeToken)) { usageLog.RateMultiplier = videoMultiplier @@ -363,10 +384,19 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( multiplier float64, imageMultiplier float64, videoMultiplier float64, + webSearchMultiplier float64, tokens UsageTokens, serviceTier string, + longContextBillingEnabled bool, ) (*CostBreakdown, error) { billingModel := firstUsageBillingModel(billingModels) + if result != nil && result.WebSearchCalls > 0 { + // Codex alpha/search 网页搜索按次计费:上游不返回 usage/token 字段,单价只取 + // 分组覆盖价(nil 时默认 0.01 = 官方 $10/1000 次),不参与渠道级模型定价。 + // 倍率与 image/video 按次口径一致:使用不含高峰因子的基础倍率 + //(用户专属 > 分组 rate_multiplier > 系统默认),与分组表单的价格预览承诺一致。 + return s.billingService.CalculateWebSearchCost(result.WebSearchCalls, webSearchPricePerCallFromAPIKey(apiKey), webSearchMultiplier), nil + } if isGrokVideoUsageResult(result, billingModels) { if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken { return s.calculateOpenAIVideoCost(ctx, billingModel, apiKey, result, videoMultiplier), nil @@ -387,7 +417,15 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( if candidate == "" { continue } - cost, err := s.calculateOpenAIRecordUsageTokenCost(ctx, apiKey, candidate, multiplier, tokens, serviceTier) + cost, err := s.calculateOpenAIRecordUsageTokenCost( + ctx, + apiKey, + candidate, + multiplier, + tokens, + serviceTier, + longContextBillingEnabled, + ) if err == nil { return cost, nil } @@ -435,21 +473,29 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( multiplier float64, tokens UsageTokens, serviceTier string, + longContextBillingEnabled bool, ) (*CostBreakdown, error) { if s.resolver != nil && apiKey.Group != nil { gid := apiKey.Group.ID return s.billingService.CalculateCostUnified(CostInput{ - Ctx: ctx, - Model: billingModel, - GroupID: &gid, - Tokens: tokens, - RequestCount: 1, - RateMultiplier: multiplier, - ServiceTier: serviceTier, - Resolver: s.resolver, + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + ServiceTier: serviceTier, + Resolver: s.resolver, + LongContextBillingEnabled: &longContextBillingEnabled, }) } - return s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) + return s.billingService.calculateCostWithServiceTierPolicy( + billingModel, + tokens, + multiplier, + serviceTier, + longContextBillingEnabled, + ) } func (s *OpenAIGatewayService) calculateOpenAIImageCost( diff --git a/backend/internal/service/openai_gpt56_max_test.go b/backend/internal/service/openai_gpt56_max_test.go index 272eb16ff0..cbca2ff3ee 100644 --- a/backend/internal/service/openai_gpt56_max_test.go +++ b/backend/internal/service/openai_gpt56_max_test.go @@ -223,13 +223,17 @@ func TestOpenAIGatewayServiceForwardOAuthCompactDowngradesMaxEffort(t *testing.T require.Equal(t, "xhigh", *result.ReasoningEffort) } -func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.T) { +func TestOpenAIGatewayServiceForwardOAuthRemoteCompactV2PreservesResponsesWire(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(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" + + "data: [DONE]\n\n", + )), }, } cfg := &config.Config{} @@ -244,6 +248,9 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing. Credentials: map[string]any{ "access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc", + "compact_model_mapping": map[string]any{ + "gpt-5.6-sol": "gpt-5.6-sol-openai-compact", + }, }, Status: StatusActive, Schedulable: true, @@ -251,16 +258,82 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing. rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) - body := []byte(`{"model":"gpt-5.6-sol","instructions":"response-test","input":"hello","reasoning":{"effort":"max"}}`) + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`) result, err := svc.Forward(context.Background(), c, account, body) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, upstream.lastReq) require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String()) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String()) require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String()) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) + require.Contains(t, rec.Body.String(), `"type":"compaction"`) + require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} + +func TestOpenAIGatewayServiceForwardAPIKeyRemoteCompactV2PreservesResponsesWire(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" + + "data: [DONE]\n\n", + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 11, + Name: "openai-apikey-responses", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com/v1", + "compact_model_mapping": map[string]any{ + "gpt-5.6-sol": "gpt-5.6-sol-openai-compact", + }, + }, + Extra: map[string]any{"use_responses_api": true}, + Status: StatusActive, + Schedulable: true, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.Equal(t, "https://example.com/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String()) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String()) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) + require.Contains(t, rec.Body.String(), `"type":"compaction"`) + require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`) require.NotNil(t, result.ReasoningEffort) require.Equal(t, "max", *result.ReasoningEffort) } diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index af0cdf669c..090948afd4 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -87,11 +87,13 @@ func TestOpenAIGatewayServiceForward_CodexImageInjectionRespectsGroupCapability( name string allowImages bool bridgeEnabled bool + responsesLite bool wantInjected bool }{ {name: "disabled group skips injection", allowImages: false, bridgeEnabled: true, wantInjected: false}, {name: "enabled group skips injection by default", allowImages: true, bridgeEnabled: false, wantInjected: false}, {name: "enabled group injects image tool when bridge enabled", allowImages: true, bridgeEnabled: true, wantInjected: true}, + {name: "responses lite skips hosted image bridge", allowImages: true, bridgeEnabled: true, responsesLite: true, wantInjected: false}, } for _, tt := range tests { @@ -106,6 +108,9 @@ func TestOpenAIGatewayServiceForward_CodexImageInjectionRespectsGroupCapability( svc := newOpenAIImageGenerationControlTestService(upstream) svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = tt.bridgeEnabled c, _ := newOpenAIImageGenerationControlTestContext(tt.allowImages, "codex_cli_rs/0.98.0") + if tt.responsesLite { + c.Request.Header.Set(responsesLiteHeader, "true") + } account := newOpenAIImageGenerationControlTestAccount() result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.4","input":"write code","stream":false}`)) @@ -115,6 +120,11 @@ func TestOpenAIGatewayServiceForward_CodexImageInjectionRespectsGroupCapability( require.NotNil(t, upstream.lastReq) hasImageTool := gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists() require.Equal(t, tt.wantInjected, hasImageTool) + expectedLiteHeader := "" + if tt.responsesLite { + expectedLiteHeader = "true" + } + require.Equal(t, expectedLiteHeader, upstream.lastReq.Header.Get(responsesLiteHeader)) instructions := gjson.GetBytes(upstream.lastBody, "instructions").String() require.Equal(t, tt.wantInjected, strings.Contains(instructions, "image_generation")) toolChoice := gjson.GetBytes(upstream.lastBody, "tool_choice") @@ -126,6 +136,24 @@ func TestOpenAIGatewayServiceForward_CodexImageInjectionRespectsGroupCapability( } } +func TestOpenAIBuildUpstreamRequestOpenAIPassthroughForwardsResponsesLiteHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.98.0") + c.Request.Header.Set(responsesLiteHeader, "true") + + svc := newOpenAIImageGenerationControlTestService(&httpUpstreamRecorder{}) + req, err := svc.buildUpstreamRequestOpenAIPassthrough( + c.Request.Context(), + c, + newOpenAIImageGenerationControlTestAccount(), + []byte(`{"model":"gpt-5.4","input":"write code"}`), + "test-token", + ) + + require.NoError(t, err) + require.Equal(t, "true", req.Header.Get(responsesLiteHeader)) +} + func TestOpenAIGatewayServiceForward_ExplicitImageToolWorksWithBridgeDisabled(t *testing.T) { gin.SetMode(gin.TestMode) @@ -283,6 +311,44 @@ func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection( require.Contains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_CodexBridgeDoesNotInjectHostedToolAlongsideImageGenNamespace(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_namespace_image","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = true + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1") + account := newOpenAIImageGenerationControlTestAccount() + body := []byte(`{ + "model":"gpt-5.5", + "stream":false, + "tools":[ + {"type":"function","name":"shell","parameters":{"type":"object"}}, + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]} + ], + "input":[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"draw a cat"}]}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]} + ], + "tool_choice":"auto" + }`) + + 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.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `tools.#(name=="image_gen").type`).String()) + require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `input.#(type=="additional_tools").tools.#(name=="image_gen").type`).String()) +} + func TestOpenAIGatewayServiceForward_CodexBridgePreservesExistingToolChoice(t *testing.T) { gin.SetMode(gin.TestMode) @@ -455,13 +521,13 @@ func TestOpenAIGatewayServiceHandleResponsesImageOutputs_Streaming(t *testing.T) gin.SetMode(gin.TestMode) svc := newOpenAIImageGenerationControlTestService(&httpUpstreamRecorder{}) - c, _ := newOpenAIImageGenerationControlTestContext(true, "unit-test-agent/1.0") + c, recorder := newOpenAIImageGenerationControlTestContext(true, "unit-test-agent/1.0") resp := &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader( - "data: {\"type\":\"response.output_item.done\",\"item\":{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"result\":\"final-image\"}}\n\n" + - "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_image_stream\",\"model\":\"gpt-5.5\",\"output\":[{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"result\":\"final-image\"}],\"usage\":{\"input_tokens\":11,\"output_tokens\":5,\"output_tokens_details\":{\"image_tokens\":4}}}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"item\":{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"status\":\"generating\",\"result\":\"final-image\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_image_stream\",\"model\":\"gpt-5.5\",\"output\":[{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"status\":\"generating\",\"result\":\"final-image\"}],\"usage\":{\"input_tokens\":11,\"output_tokens\":5,\"output_tokens_details\":{\"image_tokens\":4}}}}\n\n", )), } @@ -474,6 +540,73 @@ func TestOpenAIGatewayServiceHandleResponsesImageOutputs_Streaming(t *testing.T) require.Equal(t, 11, result.usage.InputTokens) require.Equal(t, 5, result.usage.OutputTokens) require.Equal(t, 4, result.usage.ImageOutputTokens) + require.NotContains(t, recorder.Body.String(), `"status":"generating"`) + require.Equal(t, 2, strings.Count(recorder.Body.String(), `"status":"completed"`)) +} + +func TestOpenAIGatewayServiceHandleResponsesImageOutputs_StreamingPassthrough(t *testing.T) { + gin.SetMode(gin.TestMode) + + svc := newOpenAIImageGenerationControlTestService(&httpUpstreamRecorder{}) + c, recorder := newOpenAIImageGenerationControlTestContext(true, "unit-test-agent/1.0") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_item.done\",\"item\":{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"status\":\"in_progress\",\"result\":\"final-image\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_image_stream\",\"model\":\"gpt-5.5\",\"output\":[{\"id\":\"ig_stream_1\",\"type\":\"image_generation_call\",\"status\":\"in_progress\",\"result\":\"final-image\"}],\"usage\":{\"input_tokens\":11,\"output_tokens\":5,\"output_tokens_details\":{\"image_tokens\":4}}}}\n\n", + )), + } + + result, err := svc.handleStreamingResponsePassthrough(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "gpt-5.5", "gpt-5.5") + + require.NoError(t, err) + require.NotNil(t, result) + require.NotContains(t, recorder.Body.String(), `"status":"in_progress"`) + require.Equal(t, 2, strings.Count(recorder.Body.String(), `"status":"completed"`)) +} + +func TestNormalizeCompletedImageGenerationStatus(t *testing.T) { + tests := []struct { + name string + input string + want string + wantChanged bool + }{ + { + name: "output item done with result", + input: `{"type":"response.output_item.done","item":{"type":"image_generation_call","status":"generating","result":"image-data"}}`, + want: `{"type":"response.output_item.done","item":{"type":"image_generation_call","status":"completed","result":"image-data"}}`, + wantChanged: true, + }, + { + name: "terminal response only changes completed image result", + input: `{"type":"response.completed","response":{"output":[{"type":"image_generation_call","status":"in_progress","result":"image-data"},{"type":"image_generation_call","status":"failed","result":"partial-data"}]}}`, + want: `{"type":"response.completed","response":{"output":[{"type":"image_generation_call","status":"completed","result":"image-data"},{"type":"image_generation_call","status":"failed","result":"partial-data"}]}}`, + wantChanged: true, + }, + { + name: "done item without result", + input: `{"type":"response.output_item.done","item":{"type":"image_generation_call","status":"generating"}}`, + want: `{"type":"response.output_item.done","item":{"type":"image_generation_call","status":"generating"}}`, + wantChanged: false, + }, + { + name: "non-final image event", + input: `{"type":"response.output_item.added","item":{"type":"image_generation_call","status":"generating","result":"image-data"}}`, + want: `{"type":"response.output_item.added","item":{"type":"image_generation_call","status":"generating","result":"image-data"}}`, + wantChanged: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, changed := normalizeCompletedImageGenerationStatus([]byte(tt.input)) + + require.Equal(t, tt.wantChanged, changed) + require.JSONEq(t, tt.want, string(got)) + }) + } } // TestHandleStreamingResponse_CyberPolicyCapturesRealUpstreamTokens 锁定流式 diff --git a/backend/internal/service/openai_images_json_keepalive.go b/backend/internal/service/openai_images_json_keepalive.go new file mode 100644 index 0000000000..0d8b24fb0a --- /dev/null +++ b/backend/internal/service/openai_images_json_keepalive.go @@ -0,0 +1,268 @@ +package service + +import ( + "bufio" + "errors" + "net" + "net/http" + "sync" + "time" + + "github.com/gin-gonic/gin" +) + +const openAIImagesJSONKeepaliveKey = "openai_images_json_keepalive" + +// openAIImagesJSONKeepalive keeps non-streaming Images API requests alive while +// an OAuth upstream is producing SSE internally. JSON permits leading +// whitespace, so each heartbeat remains compatible with clients expecting one +// final JSON document. +// +// Once the first heartbeat is sent, the HTTP status is committed as 200. Late +// upstream errors are still returned as an OpenAI-compatible JSON error body, +// matching the status tradeoff used by the compact SSE keepalive path. +type openAIImagesJSONKeepalive struct { + mu sync.Mutex + writer gin.ResponseWriter + started bool + stopped bool + bytes int + stop chan struct{} +} + +// StartOpenAIImagesJSONKeepalive starts whitespace heartbeats for a +// non-streaming Images request. A non-positive interval disables the feature. +func StartOpenAIImagesJSONKeepalive(c *gin.Context, interval time.Duration) func() { + if c == nil || c.Writer == nil || interval <= 0 { + return func() {} + } + originalWriter := c.Writer + k := &openAIImagesJSONKeepalive{ + writer: originalWriter, + stop: make(chan struct{}), + } + c.Set(openAIImagesJSONKeepaliveKey, k) + wrappedWriter := &openAIImagesJSONKeepaliveWriter{ResponseWriter: originalWriter, k: k} + c.Writer = wrappedWriter + + var reqDone <-chan struct{} + if c.Request != nil { + reqDone = c.Request.Context().Done() + } + go func() { + timer := time.NewTimer(interval) + defer timer.Stop() + for { + select { + case <-k.stop: + return + case <-reqDone: + return + case <-timer.C: + } + if !k.beat() { + return + } + timer.Reset(interval) + } + }() + + return func() { + k.Stop() + if current, ok := c.Writer.(*openAIImagesJSONKeepaliveWriter); ok && current == wrappedWriter { + c.Writer = originalWriter + } + } +} + +func (k *openAIImagesJSONKeepalive) beat() bool { + k.mu.Lock() + defer k.mu.Unlock() + if k.stopped { + return false + } + if !k.started { + header := k.writer.Header() + header.Set("Content-Type", "application/json; charset=utf-8") + header.Set("Cache-Control", "no-cache") + header.Set("X-Accel-Buffering", "no") + k.writer.WriteHeader(http.StatusOK) + k.started = true + } + n, err := k.writer.Write([]byte(" \n")) + k.bytes += n + if err != nil { + k.stopped = true + return false + } + k.writer.Flush() + return true +} + +func (k *openAIImagesJSONKeepalive) Stop() { + k.mu.Lock() + k.markStoppedLocked() + k.mu.Unlock() +} + +func (k *openAIImagesJSONKeepalive) markStoppedLocked() { + if k.stopped { + return + } + k.stopped = true + close(k.stop) +} + +// StopOpenAIImagesJSONKeepaliveCommitted stops heartbeats and reports whether +// they already committed a 200 response. +func StopOpenAIImagesJSONKeepaliveCommitted(c *gin.Context) bool { + k := openAIImagesJSONKeepaliveFromContext(c) + if k == nil { + return false + } + k.mu.Lock() + k.markStoppedLocked() + committed := k.started + k.mu.Unlock() + return committed +} + +// OpenAIImagesJSONKeepaliveAdjustedWrittenSize excludes heartbeat whitespace +// from response-size checks so account retry and failover remain available. +func OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c *gin.Context) int { + if c == nil || c.Writer == nil { + return -1 + } + k := openAIImagesJSONKeepaliveFromContext(c) + if k == nil { + return c.Writer.Size() + } + k.mu.Lock() + defer k.mu.Unlock() + size := k.writer.Size() + if size < 0 { + return size + } + if real := size - k.bytes; real > 0 { + return real + } + return -1 +} + +func openAIImagesJSONKeepaliveFromContext(c *gin.Context) *openAIImagesJSONKeepalive { + if c == nil { + return nil + } + value, ok := c.Get(openAIImagesJSONKeepaliveKey) + if !ok { + return nil + } + k, _ := value.(*openAIImagesJSONKeepalive) + return k +} + +type openAIImagesJSONKeepaliveWriter struct { + gin.ResponseWriter + k *openAIImagesJSONKeepalive +} + +func (w *openAIImagesJSONKeepaliveWriter) suspend() { + if w.k != nil { + w.k.Stop() + } +} + +func (w *openAIImagesJSONKeepaliveWriter) Header() http.Header { + w.suspend() + if w.ResponseWriter == nil { + return http.Header{} + } + return w.ResponseWriter.Header() +} + +func (w *openAIImagesJSONKeepaliveWriter) Write(data []byte) (int, error) { + w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } + return w.ResponseWriter.Write(data) +} + +func (w *openAIImagesJSONKeepaliveWriter) WriteString(s string) (int, error) { + w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } + return w.ResponseWriter.WriteString(s) +} + +func (w *openAIImagesJSONKeepaliveWriter) WriteHeader(code int) { + w.suspend() + if w.ResponseWriter != nil { + w.ResponseWriter.WriteHeader(code) + } +} + +func (w *openAIImagesJSONKeepaliveWriter) WriteHeaderNow() { + w.suspend() + if w.ResponseWriter != nil { + w.ResponseWriter.WriteHeaderNow() + } +} + +func (w *openAIImagesJSONKeepaliveWriter) Flush() { + w.suspend() + if w.ResponseWriter != nil { + w.ResponseWriter.Flush() + } +} + +func (w *openAIImagesJSONKeepaliveWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + if w.ResponseWriter == nil { + return nil, nil, errors.New("response writer released") + } + return w.ResponseWriter.Hijack() +} + +func (w *openAIImagesJSONKeepaliveWriter) CloseNotify() <-chan bool { + if w.ResponseWriter == nil { + ch := make(chan bool) + close(ch) + return ch + } + return w.ResponseWriter.CloseNotify() +} + +func (w *openAIImagesJSONKeepaliveWriter) Pusher() http.Pusher { + if w.ResponseWriter == nil { + return nil + } + return w.ResponseWriter.Pusher() +} + +func (w *openAIImagesJSONKeepaliveWriter) Status() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Status() +} + +func (w *openAIImagesJSONKeepaliveWriter) Size() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Size() +} + +func (w *openAIImagesJSONKeepaliveWriter) Written() bool { + if w.k == nil || w.ResponseWriter == nil { + return false + } + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Written() +} diff --git a/backend/internal/service/openai_images_json_keepalive_test.go b/backend/internal/service/openai_images_json_keepalive_test.go new file mode 100644 index 0000000000..a7207b3adb --- /dev/null +++ b/backend/internal/service/openai_images_json_keepalive_test.go @@ -0,0 +1,252 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestOpenAIImagesJSONKeepalive_PreservesValidJSONResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + originalWriter := c.Writer + + stop := StartOpenAIImagesJSONKeepalive(c, 5*time.Millisecond) + waitForOpenAIImagesJSONKeepalive(t, c) + require.Equal(t, -1, OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c)) + + c.JSON(http.StatusOK, gin.H{"data": []gin.H{{"b64_json": "aW1hZ2U="}}}) + stop() + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "application/json; charset=utf-8", rec.Header().Get("Content-Type")) + require.Equal(t, "no", rec.Header().Get("X-Accel-Buffering")) + require.True(t, rec.Flushed) + require.True(t, json.Valid(rec.Body.Bytes()), rec.Body.String()) + require.Equal(t, "aW1hZ2U=", gjson.Get(rec.Body.String(), "data.0.b64_json").String()) + require.Greater(t, OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), 0) + require.Same(t, originalWriter, c.Writer) +} + +func TestOpenAIImagesJSONKeepalive_DisabledIsNoop(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + originalWriter := c.Writer + + stop := StartOpenAIImagesJSONKeepalive(c, 0) + stop() + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"message": "invalid request"}}) + + require.Same(t, originalWriter, c.Writer) + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Equal(t, "invalid request", gjson.Get(rec.Body.String(), "error.message").String()) +} + +func TestOpenAIImagesJSONKeepalive_FastErrorPreservesStatus(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + + stop := StartOpenAIImagesJSONKeepalive(c, time.Second) + wrote := writeOpenAIImagesUpstreamErrorResponse(c, &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadRequest, + ErrorType: "invalid_request_error", + Message: "invalid size", + }) + stop() + + require.True(t, wrote) + require.Equal(t, http.StatusBadRequest, rec.Code) + require.False(t, strings.HasPrefix(rec.Body.String(), " \n")) + require.Equal(t, "invalid size", gjson.Get(rec.Body.String(), "error.message").String()) +} + +func TestOpenAIImagesJSONKeepalive_LateErrorRemainsJSON(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + + stop := StartOpenAIImagesJSONKeepalive(c, 5*time.Millisecond) + defer stop() + waitForOpenAIImagesJSONKeepalive(t, c) + + wrote := writeOpenAIImagesUpstreamErrorResponse(c, &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadRequest, + ErrorType: "image_generation_user_error", + Code: "moderation_blocked", + Message: "request rejected", + }) + + require.True(t, wrote) + require.Equal(t, http.StatusOK, rec.Code, "heartbeat already committed the status") + require.True(t, json.Valid(rec.Body.Bytes()), rec.Body.String()) + require.Equal(t, "moderation_blocked", gjson.Get(rec.Body.String(), "error.code").String()) + require.Equal(t, "request rejected", gjson.Get(rec.Body.String(), "error.message").String()) +} + +func TestOpenAIImagesJSONKeepalive_DoesNotBlockFailoverDetection(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + + stop := StartOpenAIImagesJSONKeepalive(c, 5*time.Millisecond) + waitForOpenAIImagesJSONKeepalive(t, c) + + before := OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) + require.Equal(t, -1, before) + require.True(t, c.Writer.Written()) + require.Equal(t, before, OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c)) + stop() + require.True(t, strings.TrimSpace(rec.Body.String()) == "") +} + +func TestOpenAIImagesJSONKeepalive_KeepsOAuthNonStreamResponseValid(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + + reader, writer := io.Pipe() + go func() { + time.Sleep(20 * time.Millisecond) + _, _ = io.WriteString(writer, + "data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000000,\"output\":[{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2U=\",\"output_format\":\"png\"}]}}\n\n"+ + "data: [DONE]\n\n", + ) + _ = writer.Close() + }() + + stop := StartOpenAIImagesJSONKeepalive(c, 5*time.Millisecond) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: reader, + } + svc := &OpenAIGatewayService{} + _, imageCount, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + stop() + + require.NoError(t, err) + require.Equal(t, 1, imageCount) + require.True(t, rec.Flushed) + require.True(t, strings.HasPrefix(rec.Body.String(), " \n"), rec.Body.String()) + require.True(t, json.Valid(rec.Body.Bytes()), rec.Body.String()) + require.Equal(t, "aW1hZ2U=", gjson.Get(rec.Body.String(), "data.0.b64_json").String()) +} + +func TestOpenAIImagesJSONKeepaliveWriter_NilGuards(t *testing.T) { + w := &openAIImagesJSONKeepaliveWriter{} + require.NotPanics(t, func() { + require.NotNil(t, w.Header()) + _, _ = w.Write([]byte("test")) + _, _ = w.WriteString("test") + w.WriteHeader(http.StatusOK) + w.WriteHeaderNow() + w.Flush() + require.Equal(t, 0, w.Status()) + require.Equal(t, 0, w.Size()) + require.False(t, w.Written()) + require.Nil(t, w.Pusher()) + }) + + conn, _, err := w.Hijack() + require.Error(t, err) + require.Nil(t, conn) + select { + case <-w.CloseNotify(): + default: + t.Fatal("nil writer CloseNotify channel should be closed") + } +} + +// 回归:failover 第 2+ 轮时,上一轮心跳残留的空白字节不得被误判为"已写响应", +// 可重试上游错误必须仍转换为 UpstreamFailoverError(而非裸错误吞掉换号)。 +func TestOpenAIImagesJSONKeepalive_HeartbeatBeforeForwardStillFailsOver(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat","response_format":"b64_json"}`) + + req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = req + + svc := &OpenAIGatewayService{ + httpUpstream: &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"req_img_heartbeat_failover"}, + }, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.created\",\"response\":{\"created_at\":1710000021}}\n\n" + + "data: {\"type\":\"error\",\"error\":{\"type\":\"server_error\",\"code\":\"server_error\",\"message\":\"The image service is temporarily unavailable.\"}}\n\n", + )), + }, + }, + } + parsed, err := svc.ParseOpenAIImagesRequest(c, body) + require.NoError(t, err) + + // 模拟上一轮 failover 已发生:心跳已提交 200 并写出空白字节。 + stop := StartOpenAIImagesJSONKeepalive(c, 5*time.Millisecond) + defer stop() + waitForOpenAIImagesJSONKeepalive(t, c) + + account := &Account{ + ID: 22, + Name: "openai-oauth-heartbeat-failover", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token-123", + }, + } + + result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.Contains(t, string(failoverErr.ResponseBody), "temporarily unavailable") + require.Empty(t, strings.TrimSpace(rec.Body.String()), "only heartbeat whitespace may reach the client") + + rawEvents, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 1) + require.Equal(t, "failover", events[0].Kind) + require.Equal(t, account.ID, events[0].AccountID) + require.Equal(t, http.StatusBadGateway, events[0].UpstreamStatusCode) +} + +func waitForOpenAIImagesJSONKeepalive(t *testing.T, c *gin.Context) { + t.Helper() + k := openAIImagesJSONKeepaliveFromContext(c) + require.NotNil(t, k) + require.Eventually(t, func() bool { + k.mu.Lock() + defer k.mu.Unlock() + return k.started + }, time.Second, time.Millisecond) +} diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 04347d6f90..01892ec39f 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -1010,9 +1010,13 @@ func buildOpenAIImagesStreamErrorBodyFromUpstream(err *OpenAIImagesUpstreamError } func writeOpenAIImagesUpstreamErrorResponse(c *gin.Context, err *OpenAIImagesUpstreamError) bool { - if c == nil || c.Writer == nil || c.Writer.Written() || err == nil { + if c == nil || c.Writer == nil || err == nil { return false } + if c.Writer.Written() && OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) >= 0 { + return false + } + StopOpenAIImagesJSONKeepaliveCommitted(c) errorObj := gin.H{ "type": err.clientErrorType(), "message": err.clientMessage(), @@ -1176,7 +1180,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( var sseData openAISSEDataAccumulator var processDataErr error processDataDone := false - writerSizeBeforeResponse := c.Writer.Size() + writerSizeBeforeResponse := OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) processData := func(dataBytes []byte) { if processDataDone || processDataErr != nil { @@ -1599,7 +1603,9 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( imageOutputSizes []string firstTokenMs *int ) - writerSizeBeforeResponse := c.Writer.Size() + // 与 handleOpenAIImagesOAuthResponseError 的比较端同口径:排除非流式 JSON + // keepalive 心跳字节,避免 failover 第 2 轮起把上一轮心跳残留误判为已写响应。 + writerSizeBeforeResponse := OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) if parsed.Stream { usage, imageCount, imageOutputSizes, firstTokenMs, err = s.handleOpenAIImagesOAuthStreamingResponse(resp, c, startTime, parsed.ResponseFormat, openAIImagesStreamPrefix(parsed), requestModel) if err != nil { @@ -1680,7 +1686,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthResponseError( } retryable := IsOpenAIImagesRetryableUpstreamError(upstreamErr) - responseWritten := c != nil && c.Writer != nil && c.Writer.Size() != writerSizeBeforeResponse + responseWritten := c != nil && c.Writer != nil && OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeResponse kind := "http_error" if retryable { kind = "failover" diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go index bafd36449b..db7804a4f3 100644 --- a/backend/internal/service/openai_messages_dispatch_test.go +++ b/backend/internal/service/openai_messages_dispatch_test.go @@ -37,3 +37,25 @@ func TestGroupResolveMessagesDispatchModel_GrokMapsClaudeFamilyToGrok(t *testing require.Empty(t, group.ResolveMessagesDispatchModel("grok")) require.Empty(t, group.ResolveMessagesDispatchModel("gpt-5.3-codex")) } + +func TestSanitizeGroupMessagesDispatchFields_ClearsNonOpenAIPlatform(t *testing.T) { + t.Parallel() + + group := &Group{ + Platform: PlatformAnthropic, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.6-sol", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: "gpt-5.3-codex", + ExactModelMappings: map[string]string{ + "claude-fable-5": "gpt-5.6-sol", + }, + }, + } + + sanitizeGroupMessagesDispatchFields(group) + + require.False(t, group.AllowMessagesDispatch) + require.Empty(t, group.DefaultMappedModel) + require.Equal(t, OpenAIMessagesDispatchModelConfig{}, group.MessagesDispatchModelConfig) +} diff --git a/backend/internal/service/openai_model_mapping.go b/backend/internal/service/openai_model_mapping.go index cb7a8ca84b..8ba1d6fe1b 100644 --- a/backend/internal/service/openai_model_mapping.go +++ b/backend/internal/service/openai_model_mapping.go @@ -3,19 +3,20 @@ package service import "strings" // resolveOpenAIForwardModel 解析 OpenAI 兼容转发使用的模型。 -// defaultMappedModel 只服务于 /v1/messages 的 Claude 系列显式调度映射, -// 不作为普通 OpenAI 请求的未知模型兜底。 -func resolveOpenAIForwardModel(account *Account, requestedModel, defaultMappedModel string) string { +// messagesDispatchMappedModel 是调用方已为 /v1/messages 解析的显式调度结果; +// 普通 OpenAI 请求必须传空,避免将分组配置作为通用模型兜底。 +func resolveOpenAIForwardModel(account *Account, requestedModel, messagesDispatchMappedModel string) string { + messagesDispatchMappedModel = strings.TrimSpace(messagesDispatchMappedModel) if account == nil { - if defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" { - return defaultMappedModel + if messagesDispatchMappedModel != "" { + return messagesDispatchMappedModel } return requestedModel } mappedModel, matched := account.ResolveMappedModel(requestedModel) - if !matched && defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" { - return defaultMappedModel + if !matched && messagesDispatchMappedModel != "" { + return messagesDispatchMappedModel } return mappedModel } diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index f2ceb3551c..7107a706ad 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -4,159 +4,156 @@ import "testing" func TestResolveOpenAIForwardModel(t *testing.T) { tests := []struct { - name string - account *Account - requestedModel string - defaultMappedModel string - expectedModel string + name string + account *Account + requestedModel string + messagesDispatchMappedModel string + expectedModel string }{ { - name: "uses messages dispatch default for claude model", + name: "uses messages dispatch model for known claude family", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "claude-opus-4-6", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-4o-mini", + requestedModel: "claude-opus-4-6", + messagesDispatchMappedModel: "gpt-4o-mini", + expectedModel: "gpt-4o-mini", }, { - name: "does not fall back to group default for invalid gpt model", + name: "uses exact messages dispatch model for unknown claude family", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt6", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt6", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: " gpt-5.6-sol ", + expectedModel: "gpt-5.6-sol", }, { - name: "preserves explicit gpt-5.4 instead of group default", + name: "nil account uses messages dispatch model", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.6-sol", + }, + { + name: "nil account without messages dispatch keeps requested model", + requestedModel: "claude-fable-5", + expectedModel: "claude-fable-5", + }, + { + name: "ordinary unknown gpt model has no messages dispatch fallback", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "gpt6", + expectedModel: "gpt6", }, { - name: "preserves exact passthrough mapping instead of group default", + name: "account exact mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-5.4": "gpt-5.4", + "claude-fable-5": "gpt-5.5", }, }, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.5", }, { - name: "preserves wildcard passthrough mapping instead of group default", + name: "account wildcard mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-*": "gpt-5.4", + "claude-*": "gpt-5.4", }, }, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.4", }, { - name: "uses account remap when explicit target differs", + name: "account passthrough mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-5": "gpt-5.4", + "claude-fable-5": "claude-fable-5", }, }, }, - requestedModel: "gpt-5", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "claude-fable-5", }, { - name: "preserves codex spark instead of group default", + name: "ordinary codex spark request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.3-codex-spark", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.3-codex-spark", + requestedModel: "gpt-5.3-codex-spark", + expectedModel: "gpt-5.3-codex-spark", }, { - name: "preserves gpt-5.5 instead of group default", + name: "ordinary gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.5", + requestedModel: "gpt-5.5", + expectedModel: "gpt-5.5", }, { - name: "preserves gpt-5.5-pro instead of group default", + name: "ordinary gpt-5.5-pro request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5-pro", - defaultMappedModel: "gpt-5.5", - expectedModel: "gpt-5.5-pro", + requestedModel: "gpt-5.5-pro", + expectedModel: "gpt-5.5-pro", }, { - name: "preserves compact-spelled gpt5.5 instead of group default", + name: "ordinary compact-spelled gpt5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt5.5", + requestedModel: "gpt5.5", + expectedModel: "gpt5.5", }, { - name: "preserves openai namespaced gpt-5.5 instead of group default", + name: "ordinary namespaced gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "openai/gpt-5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "openai/gpt-5.5", + requestedModel: "openai/gpt-5.5", + expectedModel: "openai/gpt-5.5", }, { - name: "preserves compact gpt-5.5 instead of group default", + name: "ordinary compact gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5-openai-compact", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.5-openai-compact", + requestedModel: "gpt-5.5-openai-compact", + expectedModel: "gpt-5.5-openai-compact", + }, + { + name: "whitespace-only messages dispatch model is ignored", + account: &Account{ + Credentials: map[string]any{}, + }, + requestedModel: "gpt-5.5", + messagesDispatchMappedModel: " ", + expectedModel: "gpt-5.5", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.defaultMappedModel); got != tt.expectedModel { + if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.messagesDispatchMappedModel); got != tt.expectedModel { t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", got, tt.expectedModel) } }) } } -func TestResolveOpenAIForwardModel_PreventsClaudeModelFromFallingBackToGpt54(t *testing.T) { - account := &Account{ - Credentials: map[string]any{}, - } - - withoutDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "") - if withoutDefault != "claude-opus-4-6" { - t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withoutDefault, "claude-opus-4-6") - } - - withDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "gpt-5.4") - if withDefault != "gpt-5.4" { - t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withDefault, "gpt-5.4") - } -} - func TestResolveOpenAICompactForwardModel(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index de8ecf030b..0aa536f94f 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -14,6 +14,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" @@ -347,6 +348,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali c.Request.Header.Set("Accept-Encoding", "gzip") c.Request.Header.Set("Proxy-Authorization", "Basic abc") c.Request.Header.Set("X-Test", "keep") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") originalBody := []byte(`{"model":"gpt-5.2","stream":true,"store":true,"instructions":"local-test-instructions","input":[{"type":"text","text":"hi"}]}`) @@ -409,6 +411,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali require.Empty(t, upstream.lastReq.Header.Get("Accept-Encoding")) require.Empty(t, upstream.lastReq.Header.Get("Proxy-Authorization")) require.Empty(t, upstream.lastReq.Header.Get("X-Test")) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) // 3) required OAuth headers are present require.Equal(t, "chatgpt.com", upstream.lastReq.Host) @@ -420,6 +423,194 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali require.NotContains(t, body, "\"name\":\"edit\"") } +func TestOpenAIGatewayService_OAuthPassthrough_NamespaceRequestAndStreamResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + + originalBody := []byte(`{ + "model":"gpt-5.5", + "stream":true, + "instructions":"local-test-instructions", + "tools":[ + {"type":"function","name":"plain","description":"keep","parameters":{"type":"object"}}, + {"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","description":"spawn","parameters":{"type":"object"}}]} + ], + "tool_choice":{"type":"function","name":"spawn_agent","namespace":"collaboration"}, + "input":[{"type":"function_call","call_id":"call_old","name":"spawn_agent","namespace":"collaboration","arguments":"{}"}] + }`) + + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"collaboration__spawn_agent","arguments":""}}`, + "", + `data: {"type":"response.output_item.done","output_index":0,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"collaboration__spawn_agent","arguments":"{}"}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_1","status":"completed","output":[{"type":"function_call","id":"fc_1","call_id":"call_1","name":"collaboration__spawn_agent","arguments":"{}"}],"usage":{"input_tokens":2,"output_tokens":1,"total_tokens":3}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_namespace"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 123, Name: "acc", Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"}, + Extra: map[string]any{"openai_passthrough": true}, Status: StatusActive, Schedulable: true, RateMultiplier: f64p(1), + } + + result, err := svc.Forward(context.Background(), c, account, originalBody) + require.NoError(t, err) + require.NotNil(t, result) + + require.Len(t, gjson.GetBytes(upstream.lastBody, "tools").Array(), 2) + require.Equal(t, "plain", gjson.GetBytes(upstream.lastBody, "tools.0.name").String()) + require.Equal(t, "function", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) + require.Equal(t, "collaboration__spawn_agent", gjson.GetBytes(upstream.lastBody, "tools.1.name").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "tools.1.tools").Exists()) + require.Equal(t, "collaboration__spawn_agent", gjson.GetBytes(upstream.lastBody, "tool_choice.name").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice.namespace").Exists()) + require.Equal(t, "collaboration__spawn_agent", gjson.GetBytes(upstream.lastBody, "input.0.name").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.namespace").Exists()) + + downstream := rec.Body.String() + require.NotContains(t, downstream, "collaboration__spawn_agent") + require.Contains(t, downstream, `"name":"spawn_agent"`) + require.Contains(t, downstream, `"namespace":"collaboration"`) +} + +func TestOpenAIGatewayService_NativeOAuth_NamespaceRequestAndStreamResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + body := []byte(`{ + "model":"gpt-5.5","stream":true,"instructions":"test", + "tools":[{"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","parameters":{"type":"object"}}]}], + "input":[{"type":"function_call","call_id":"call_old","name":"spawn_agent","namespace":"collaboration","arguments":"{}"}] + }`) + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"collaboration__spawn_agent","arguments":""}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_1","status":"completed","output":[{"type":"function_call","id":"fc_1","call_id":"call_1","name":"collaboration__spawn_agent","arguments":"{}"}],"usage":{"input_tokens":2,"output_tokens":1,"total_tokens":3}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_native_namespace"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 124, Name: "native", Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"}, + Status: StatusActive, Schedulable: true, RateMultiplier: f64p(1), + } + + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "function", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "collaboration__spawn_agent", gjson.GetBytes(upstream.lastBody, "tools.0.name").String()) + require.Equal(t, "collaboration__spawn_agent", gjson.GetBytes(upstream.lastBody, "input.0.name").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.namespace").Exists()) + require.NotContains(t, rec.Body.String(), "collaboration__spawn_agent") + require.Contains(t, rec.Body.String(), `"name":"spawn_agent"`) + require.Contains(t, rec.Body.String(), `"namespace":"collaboration"`) +} + +func TestOpenAIGatewayService_NativeOAuth_NamespaceNonStreamingResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + setOpenAIResponsesNamespaceNames(c, map[string]apicompat.ResponsesNamespaceName{ + "collaboration__spawn_agent": {Namespace: "collaboration", Name: "spawn_agent"}, + }) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "id":"resp_1","output":[{"type":"function_call","name":"collaboration__spawn_agent","call_id":"call_1","arguments":"{}"}], + "usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2} + }`)), + } + + result, err := (&OpenAIGatewayService{cfg: &config.Config{}}).handleNonStreamingResponse( + context.Background(), resp, c, &Account{Type: AccountTypeOAuth}, "gpt-5.5", "gpt-5.5", + ) + require.NoError(t, err) + require.NotNil(t, result) + require.NotContains(t, rec.Body.String(), "collaboration__spawn_agent") + require.Contains(t, rec.Body.String(), `"name":"spawn_agent"`) + require.Contains(t, rec.Body.String(), `"namespace":"collaboration"`) +} + +func TestOpenAIGatewayService_OAuthPassthrough_NamespaceNonStreamingResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_1","output":[{"type":"function_call","name":"collaboration__spawn_agent","call_id":"call_1","arguments":"{}"}],"usage":{"input_tokens":1,"output_tokens":1}}`)), + } + names := map[string]apicompat.ResponsesNamespaceName{ + "collaboration__spawn_agent": {Namespace: "collaboration", Name: "spawn_agent"}, + } + setOpenAIResponsesNamespaceNames(c, names) + + result, err := (&OpenAIGatewayService{cfg: &config.Config{}}).handleNonStreamingResponsePassthrough( + context.Background(), resp, c, "gpt-5.5", "", + ) + require.NoError(t, err) + require.NotNil(t, result) + require.NotContains(t, rec.Body.String(), "collaboration__spawn_agent") + require.Contains(t, rec.Body.String(), `"name":"spawn_agent"`) + require.Contains(t, rec.Body.String(), `"namespace":"collaboration"`) +} + +func TestOpenAIGatewayService_OAuthPassthrough_NamespaceCollisionReturnsBadRequest(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + body := []byte(`{ + "model":"gpt-5.5","stream":true,"instructions":"test", + "tools":[ + {"type":"function","name":"collaboration__spawn_agent","parameters":{"type":"object"}}, + {"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","parameters":{"type":"object"}}]} + ],"input":"hi" + }`) + upstream := &httpUpstreamRecorder{} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 123, Name: "acc", Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"}, + Extra: map[string]any{"openai_passthrough": true}, Status: StatusActive, Schedulable: true, RateMultiplier: f64p(1), + } + + result, err := svc.Forward(context.Background(), c, account, body) + require.Error(t, err) + require.Nil(t, result) + require.Nil(t, upstream.lastReq) + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Equal(t, "invalid_request_error", gjson.Get(rec.Body.String(), "error.type").String()) + require.Equal(t, "tools", gjson.Get(rec.Body.String(), "error.param").String()) + require.Contains(t, gjson.Get(rec.Body.String(), "error.message").String(), "conflicts with a top-level tool") +} + func TestOpenAIGatewayService_OAuthPassthrough_CompactUsesJSONAndKeepsNonStreaming(t *testing.T) { gin.SetMode(gin.TestMode) @@ -1373,6 +1564,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) c.Request.Header.Set("User-Agent", "curl/8.0") c.Request.Header.Set("X-Test", "keep") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") originalBody := []byte(`{"model":"gpt-5.2","stream":false,"service_tier":"flex","max_output_tokens":128,"input":[{"type":"text","text":"hi"}]}`) resp := &http.Response{ @@ -1410,6 +1602,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd require.Equal(t, "https://api.openai.com/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer sk-api-key", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "curl/8.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) require.Empty(t, upstream.lastReq.Header.Get("X-Test")) } diff --git a/backend/internal/service/openai_quota_reset_credits.go b/backend/internal/service/openai_quota_reset_credits.go new file mode 100644 index 0000000000..75756976db --- /dev/null +++ b/backend/internal/service/openai_quota_reset_credits.go @@ -0,0 +1,141 @@ +package service + +import ( + "bytes" + "encoding/json" + "strconv" + "strings" +) + +type openAIRateLimitResetCreditDetailPayload struct { + ExpiresAt string `json:"expires_at,omitempty"` + ExpiresAtCamel string `json:"expiresAt,omitempty"` + ResetType string `json:"reset_type,omitempty"` + ResetTypeCamel string `json:"resetType,omitempty"` + Status string `json:"status,omitempty"` +} + +type openAIRateLimitResetCreditDetailsPayload struct { + AvailableCount json.RawMessage `json:"available_count,omitempty"` + AvailableCountCamel json.RawMessage `json:"availableCount,omitempty"` + Credits json.RawMessage `json:"credits,omitempty"` + RateLimitResetCredits json.RawMessage `json:"rate_limit_reset_credits,omitempty"` + Items json.RawMessage `json:"items,omitempty"` + Data json.RawMessage `json:"data,omitempty"` +} + +type openAIRateLimitResetCreditDetails struct { + AvailableCount *int + AvailableCreditCount int + CreditListPresent bool + Credits []OpenAIRateLimitResetCreditDetail +} + +func parseOpenAIRateLimitResetCreditDetails(body []byte) (openAIRateLimitResetCreditDetails, error) { + trimmed := bytes.TrimSpace(body) + if len(trimmed) == 0 { + return openAIRateLimitResetCreditDetails{}, nil + } + + var rawCredits []*openAIRateLimitResetCreditDetailPayload + var availableCount *int + var creditListPresent bool + if trimmed[0] == '[' { + if err := json.Unmarshal(trimmed, &rawCredits); err != nil { + return openAIRateLimitResetCreditDetails{}, err + } + creditListPresent = true + } else { + var payload openAIRateLimitResetCreditDetailsPayload + if err := json.Unmarshal(trimmed, &payload); err != nil { + return openAIRateLimitResetCreditDetails{}, err + } + availableCount = parseOpenAIResetCreditAvailableCount(payload.AvailableCount, payload.AvailableCountCamel) + var err error + rawCredits, creditListPresent, err = firstPresentResetCreditPayload( + payload.Credits, + payload.RateLimitResetCredits, + payload.Items, + payload.Data, + ) + if err != nil { + return openAIRateLimitResetCreditDetails{}, err + } + } + + credits := make([]OpenAIRateLimitResetCreditDetail, 0, len(rawCredits)) + availableCreditCount := 0 + for _, raw := range rawCredits { + if raw == nil { + continue + } + resetType := strings.TrimSpace(raw.ResetType) + if resetType == "" { + resetType = strings.TrimSpace(raw.ResetTypeCamel) + } + if resetType != "" && !strings.EqualFold(resetType, "codex_rate_limits") { + continue + } + if status := strings.TrimSpace(raw.Status); status != "" && !strings.EqualFold(status, "available") { + continue + } + availableCreditCount++ + expiresAt := strings.TrimSpace(raw.ExpiresAt) + if expiresAt == "" { + expiresAt = strings.TrimSpace(raw.ExpiresAtCamel) + } + if expiresAt == "" { + continue + } + credits = append(credits, OpenAIRateLimitResetCreditDetail{ExpiresAt: expiresAt}) + } + return openAIRateLimitResetCreditDetails{ + AvailableCount: availableCount, + AvailableCreditCount: availableCreditCount, + CreditListPresent: creditListPresent, + Credits: credits, + }, nil +} + +func parseOpenAIResetCreditAvailableCount(values ...json.RawMessage) *int { + for _, value := range values { + trimmed := bytes.TrimSpace(value) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + continue + } + + var count int + if trimmed[0] == '"' { + var text string + if err := json.Unmarshal(trimmed, &text); err != nil { + continue + } + parsed, err := strconv.Atoi(strings.TrimSpace(text)) + if err != nil { + continue + } + count = parsed + } else if err := json.Unmarshal(trimmed, &count); err != nil { + continue + } + if count >= 0 { + return &count + } + } + return nil +} + +func firstPresentResetCreditPayload(values ...json.RawMessage) ([]*openAIRateLimitResetCreditDetailPayload, bool, error) { + for _, value := range values { + trimmed := bytes.TrimSpace(value) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + continue + } + var credits []*openAIRateLimitResetCreditDetailPayload + if err := json.Unmarshal(trimmed, &credits); err != nil { + return nil, false, err + } + return credits, true, nil + } + return nil, false, nil +} diff --git a/backend/internal/service/openai_quota_reset_credits_test.go b/backend/internal/service/openai_quota_reset_credits_test.go new file mode 100644 index 0000000000..5d994602e1 --- /dev/null +++ b/backend/internal/service/openai_quota_reset_credits_test.go @@ -0,0 +1,175 @@ +package service + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestParseOpenAIRateLimitResetCreditDetails_PreservesAvailableCreditOrder(t *testing.T) { + body := []byte(`{ + "availableCount":"2", + "credits":[ + {"reset_type":"codex_rate_limits","status":"redeemed","expires_at":"2026-07-01T04:05:06Z"}, + {"reset_type":"codex_rate_limits","status":"available","expires_at":"2026-07-04T04:05:06Z"}, + {"resetType":"codex_rate_limits","status":"available","expiresAt":"2026-07-03T04:05:06Z"}, + {"reset_type":"other","status":"available","expires_at":"2026-07-02T04:05:06Z"} + ] + }`) + + details, err := parseOpenAIRateLimitResetCreditDetails(body) + require.NoError(t, err) + require.NotNil(t, details.AvailableCount) + require.Equal(t, 2, *details.AvailableCount) + require.Equal(t, []OpenAIRateLimitResetCreditDetail{ + {ExpiresAt: "2026-07-04T04:05:06Z"}, + {ExpiresAt: "2026-07-03T04:05:06Z"}, + }, details.Credits) +} + +func TestQueryUsageResetCreditCountPrecedence(t *testing.T) { + tests := []struct { + name string + usageBody string + detailBody string + wantCount int + wantCredits int + wantNil bool + }{ + { + name: "detail count creates missing usage credits", + usageBody: `{}`, + detailBody: `{"available_count":3,"credits":[{"expires_at":"2026-07-03T04:05:06Z"}]}`, + wantCount: 3, wantCredits: 1, + }, + { + name: "explicit detail zero overrides usage and records", + usageBody: `{"rate_limit_reset_credits":{"available_count":4}}`, + detailBody: `{"available_count":0,"credits":[{"expires_at":"2026-07-03T04:05:06Z"}]}`, + wantCount: 0, wantCredits: 1, + }, + { + name: "available records override usage when detail count is absent", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `{"credits":[{"expires_at":"2026-07-03T04:05:06Z"},{"expiresAt":"2026-07-04T04:05:06Z"}]}`, + wantCount: 2, wantCredits: 2, + }, + { + name: "empty detail list overrides usage with zero", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `{"credits":[]}`, + wantCount: 0, + }, + { + name: "fully filtered list overrides usage with zero", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `{"credits":[{"reset_type":"codex_rate_limits","status":"redeemed","expires_at":"2026-07-03T04:05:06Z"},{"reset_type":"other","status":"available","expires_at":"2026-07-04T04:05:06Z"}]}`, + wantCount: 0, + }, + { + name: "available records without expiry still count", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `{"credits":[{"status":"available"},{"status":"available","expires_at":"2026-07-04T04:05:06Z"}]}`, + wantCount: 2, wantCredits: 1, + }, + { + name: "shape without count or list preserves usage details", + usageBody: `{"rate_limit_reset_credits":{"available_count":5,"credits":[{"expires_at":"usage-expiry"}]}}`, + detailBody: `{}`, + wantCount: 5, + wantCredits: 1, + }, + { + name: "negative detail count without list preserves usage", + usageBody: `{"rate_limit_reset_credits":{"available_count":4}}`, + detailBody: `{"available_count":-1}`, + wantCount: 4, + }, + { + name: "negative detail count falls back to available records", + usageBody: `{"rate_limit_reset_credits":{"available_count":4}}`, + detailBody: `{"available_count":-1,"credits":[{"status":"available","expires_at":"2026-07-04T04:05:06Z"}]}`, + wantCount: 1, wantCredits: 1, + }, + { + name: "empty object preserves missing usage credits", + usageBody: `{}`, + detailBody: `{}`, + wantNil: true, + }, + { + name: "null body preserves missing usage credits", + usageBody: `{}`, + detailBody: `null`, + wantNil: true, + }, + { + name: "empty body preserves missing usage credits", + usageBody: `{}`, + detailBody: ``, + wantNil: true, + }, + { + name: "null object record is not counted", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `{"credits":[null]}`, + wantCount: 0, + }, + { + name: "null top level record is not counted", + usageBody: `{"rate_limit_reset_credits":{"available_count":7}}`, + detailBody: `[null]`, + wantCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := &Account{ + ID: 100, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Credentials: map[string]any{ + "chatgpt_account_id": "org-parent123", + }, + } + repo := &stubQuotaAccountRepo{accounts: map[int64]*Account{100: account}} + tokenCache := &stubQuotaTokenCache{tokens: map[string]string{ + OpenAITokenCacheKey(account): "fake-token", + }} + tokenProvider := NewOpenAITokenProvider(repo, tokenCache, nil) + + var detailCalls int + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("content-type", "application/json") + switch r.URL.Path { + case "/backend-api/wham/usage": + _, _ = w.Write([]byte(tt.usageBody)) + case "/backend-api/wham/rate-limit-reset-credits": + detailCalls++ + _, _ = w.Write([]byte(tt.detailBody)) + default: + http.NotFound(w, r) + } + })) + defer srv.Close() + + svc := NewOpenAIQuotaService(repo, nil, tokenProvider, newQuotaRedirectingFactory(srv)) + usage, err := svc.QueryUsage(context.Background(), 100) + require.NoError(t, err) + require.NotNil(t, usage) + require.Equal(t, 1, detailCalls) + if tt.wantNil { + require.Nil(t, usage.RateLimitResetCredits) + return + } + require.NotNil(t, usage.RateLimitResetCredits) + require.Equal(t, tt.wantCount, usage.RateLimitResetCredits.AvailableCount) + require.Len(t, usage.RateLimitResetCredits.Credits, tt.wantCredits) + }) + } +} diff --git a/backend/internal/service/openai_quota_service.go b/backend/internal/service/openai_quota_service.go index 735a580570..3669154b64 100644 --- a/backend/internal/service/openai_quota_service.go +++ b/backend/internal/service/openai_quota_service.go @@ -1,11 +1,9 @@ package service import ( - "bytes" "context" "crypto/rand" "encoding/hex" - "encoding/json" "fmt" "log/slog" "net/http" @@ -191,13 +189,26 @@ func (s *OpenAIQuotaService) QueryUsage(ctx context.Context, accountID int64) (* } payload.FetchedAt = time.Now().Unix() - if payload.RateLimitResetCredits != nil && payload.RateLimitResetCredits.AvailableCount > 0 { - payload.RateLimitResetCredits.Credits = s.queryResetCreditDetails(callCtx, client, accessToken, chatGPTAccountID, fedRAMP, accountID) + details := s.queryResetCreditDetails(callCtx, client, accessToken, chatGPTAccountID, fedRAMP, accountID) + if details != nil { + hasDetailCount := details.AvailableCount != nil + if payload.RateLimitResetCredits == nil { + payload.RateLimitResetCredits = &OpenAIRateLimitResetCredits{} + } + if details.CreditListPresent { + payload.RateLimitResetCredits.Credits = details.Credits + } + switch { + case hasDetailCount: + payload.RateLimitResetCredits.AvailableCount = *details.AvailableCount + case details.CreditListPresent: + payload.RateLimitResetCredits.AvailableCount = details.AvailableCreditCount + } } return &payload, nil } -func (s *OpenAIQuotaService) queryResetCreditDetails(ctx context.Context, client *req.Client, accessToken, chatGPTAccountID string, fedRAMP bool, accountID int64) []OpenAIRateLimitResetCreditDetail { +func (s *OpenAIQuotaService) queryResetCreditDetails(ctx context.Context, client *req.Client, accessToken, chatGPTAccountID string, fedRAMP bool, accountID int64) *openAIRateLimitResetCreditDetails { quotaHeaders, headerErr := s.buildCodexQuotaHeaders(ctx, accountID, accessToken, chatGPTAccountID, fedRAMP) if headerErr != nil { slog.Warn("openai_quota_reset_credit_details_auth_failed", "account_id", accountID, "error", headerErr) @@ -216,12 +227,15 @@ func (s *OpenAIQuotaService) queryResetCreditDetails(ctx context.Context, client return nil } - credits, err := parseOpenAIRateLimitResetCreditDetails(resp.Bytes()) + details, err := parseOpenAIRateLimitResetCreditDetails(resp.Bytes()) if err != nil { slog.Warn("openai_quota_reset_credit_details_parse_failed", "account_id", accountID, "error", err) return nil } - return credits + if details.AvailableCount == nil && !details.CreditListPresent { + return nil + } + return &details } // ResetCredit consumes one rate_limit_reset_credit for the given OpenAI account. @@ -494,65 +508,6 @@ func generateRedeemRequestID() (string, error) { return fmt.Sprintf("%s-%s-%s-%s-%s", hexStr[0:8], hexStr[8:12], hexStr[12:16], hexStr[16:20], hexStr[20:]), nil } -type openAIRateLimitResetCreditDetailPayload struct { - ExpiresAt string `json:"expires_at,omitempty"` - ExpiresAtCamel string `json:"expiresAt,omitempty"` -} - -type openAIRateLimitResetCreditDetailsPayload struct { - Credits []openAIRateLimitResetCreditDetailPayload `json:"credits,omitempty"` - RateLimitResetCredits []openAIRateLimitResetCreditDetailPayload `json:"rate_limit_reset_credits,omitempty"` - Items []openAIRateLimitResetCreditDetailPayload `json:"items,omitempty"` - Data []openAIRateLimitResetCreditDetailPayload `json:"data,omitempty"` -} - -func parseOpenAIRateLimitResetCreditDetails(body []byte) ([]OpenAIRateLimitResetCreditDetail, error) { - trimmed := bytes.TrimSpace(body) - if len(trimmed) == 0 { - return nil, nil - } - - var rawCredits []openAIRateLimitResetCreditDetailPayload - if trimmed[0] == '[' { - if err := json.Unmarshal(trimmed, &rawCredits); err != nil { - return nil, err - } - } else { - var payload openAIRateLimitResetCreditDetailsPayload - if err := json.Unmarshal(trimmed, &payload); err != nil { - return nil, err - } - rawCredits = firstNonEmptyResetCreditPayload( - payload.Credits, - payload.RateLimitResetCredits, - payload.Items, - payload.Data, - ) - } - - credits := make([]OpenAIRateLimitResetCreditDetail, 0, len(rawCredits)) - for _, raw := range rawCredits { - expiresAt := strings.TrimSpace(raw.ExpiresAt) - if expiresAt == "" { - expiresAt = strings.TrimSpace(raw.ExpiresAtCamel) - } - if expiresAt == "" { - continue - } - credits = append(credits, OpenAIRateLimitResetCreditDetail{ExpiresAt: expiresAt}) - } - return credits, nil -} - -func firstNonEmptyResetCreditPayload(lists ...[]openAIRateLimitResetCreditDetailPayload) []openAIRateLimitResetCreditDetailPayload { - for _, list := range lists { - if len(list) > 0 { - return list - } - } - return nil -} - // buildCodexSparkWindowExtraUpdates extracts Codex Spark usage windows from the // /wham/usage response body's additional_rate_limits, matching the entry with // MeteredFeature == "codex_bengalfox". It produces plain codex_* keys (NOT the diff --git a/backend/internal/service/openai_quota_spark_window_test.go b/backend/internal/service/openai_quota_spark_window_test.go index 4e5fbc4a84..c521b149e9 100644 --- a/backend/internal/service/openai_quota_spark_window_test.go +++ b/backend/internal/service/openai_quota_spark_window_test.go @@ -310,6 +310,10 @@ func TestQueryUsageAgentIdentityRecoversInvalidTaskOnce(t *testing.T) { _, _ = w.Write([]byte(`{"task_id":"task-quota-new"}`)) return } + if strings.Contains(r.URL.Path, "rate-limit-reset-credits") { + _, _ = w.Write([]byte(`{}`)) + return + } usageCalls++ if usageCalls == 1 { w.WriteHeader(http.StatusUnauthorized) @@ -372,11 +376,11 @@ func TestParseOpenAIRateLimitResetCreditDetails_CompatibleContainers(t *testing. t.Run(tt.name, func(t *testing.T) { got, err := parseOpenAIRateLimitResetCreditDetails([]byte(tt.body)) require.NoError(t, err) - require.Len(t, got, len(tt.want)) + require.Len(t, got.Credits, len(tt.want)) for i := range tt.want { - require.Equal(t, tt.want[i], got[i].ExpiresAt) + require.Equal(t, tt.want[i], got.Credits[i].ExpiresAt) } - encoded, err := json.Marshal(got) + encoded, err := json.Marshal(got.Credits) require.NoError(t, err) require.NotContains(t, string(encoded), "secret-id") }) diff --git a/backend/internal/service/openai_responses_namespace.go b/backend/internal/service/openai_responses_namespace.go new file mode 100644 index 0000000000..7d71b64814 --- /dev/null +++ b/backend/internal/service/openai_responses_namespace.go @@ -0,0 +1,83 @@ +package service + +import ( + "bytes" + "encoding/json" + "fmt" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" +) + +const openAIResponsesNamespaceNamesContextKey = "openai_responses_namespace_names" + +// shouldFlattenOpenAIResponsesNamespaces 判定原生 Responses 转发前是否摊平 +// Codex namespace 工具。WSv2 上游原生支持 namespace,且 WS 出口 +// (openai_ws_forwarder_v2)原样转发上游事件、不经 HTTP 回程还原,摊平后的 +// 平名无法还原会破坏客户端工具匹配,因此实际走 WSv2 分支的请求保持 namespace +// 原样。透传账号先于 WSv2 分支经 HTTP 转发返回,仍需摊平。 +func shouldFlattenOpenAIResponsesNamespaces(account *Account, transport OpenAIUpstreamTransport, passthroughEnabled bool) bool { + if account == nil || account.Type != AccountTypeOAuth { + return false + } + if transport == OpenAIUpstreamTransportResponsesWebsocketV2 && !passthroughEnabled { + return false + } + return true +} + +func flattenOpenAIResponsesNamespaces(c *gin.Context, body []byte) ([]byte, error) { + if !bytes.Contains(body, []byte(`"namespace"`)) { + return body, nil + } + var requestBody map[string]any + if err := json.Unmarshal(body, &requestBody); err != nil { + return body, fmt.Errorf("decode OpenAI namespace body: %w", err) + } + names, changed, err := apicompat.FlattenResponsesNamespacesExcept(requestBody, map[string]bool{"image_gen": true}) + if err != nil { + return body, err + } + if !changed { + return body, nil + } + rebuilt, err := marshalOpenAIUpstreamJSON(requestBody) + if err != nil { + return body, fmt.Errorf("encode OpenAI namespace body: %w", err) + } + setOpenAIResponsesNamespaceNames(c, names) + return rebuilt, nil +} + +func setOpenAIResponsesNamespaceNames(c *gin.Context, names map[string]apicompat.ResponsesNamespaceName) { + if c != nil && len(names) > 0 { + c.Set(openAIResponsesNamespaceNamesContextKey, names) + } +} + +func openAIResponsesNamespaceNames(c *gin.Context) map[string]apicompat.ResponsesNamespaceName { + if c == nil { + return nil + } + value, ok := c.Get(openAIResponsesNamespaceNamesContextKey) + if !ok { + return nil + } + names, _ := value.(map[string]apicompat.ResponsesNamespaceName) + return names +} + +func restoreOpenAIResponsesNamespacePayload(c *gin.Context, payload []byte) ([]byte, error) { + names := openAIResponsesNamespaceNames(c) + if len(names) == 0 || !json.Valid(payload) { + return payload, nil + } + restored, changed, err := apicompat.RestoreResponsesNamespaceCalls(payload, names) + if err != nil { + return payload, err + } + if changed { + return restored, nil + } + return payload, nil +} diff --git a/backend/internal/service/openai_responses_namespace_test.go b/backend/internal/service/openai_responses_namespace_test.go new file mode 100644 index 0000000000..7aa1392260 --- /dev/null +++ b/backend/internal/service/openai_responses_namespace_test.go @@ -0,0 +1,34 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestShouldFlattenOpenAIResponsesNamespaces(t *testing.T) { + oauth := &Account{Type: AccountTypeOAuth} + apiKey := &Account{Type: AccountTypeAPIKey} + + tests := []struct { + name string + account *Account + transport OpenAIUpstreamTransport + passthroughEnabled bool + want bool + }{ + {name: "oauth_http", account: oauth, transport: OpenAIUpstreamTransportHTTPSSE, want: true}, + {name: "oauth_http_passthrough", account: oauth, transport: OpenAIUpstreamTransportHTTPSSE, passthroughEnabled: true, want: true}, + // WSv2 出口原样转发上游事件、不做回程还原,摊平会让客户端收到无法匹配的平名。 + {name: "oauth_wsv2", account: oauth, transport: OpenAIUpstreamTransportResponsesWebsocketV2, want: false}, + // 透传账号先于 WSv2 分支经 HTTP 转发返回,仍需摊平。 + {name: "oauth_wsv2_passthrough", account: oauth, transport: OpenAIUpstreamTransportResponsesWebsocketV2, passthroughEnabled: true, want: true}, + {name: "apikey_http", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, want: false}, + {name: "nil_account", account: nil, transport: OpenAIUpstreamTransportHTTPSSE, want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, shouldFlattenOpenAIResponsesNamespaces(tt.account, tt.transport, tt.passthroughEnabled)) + }) + } +} diff --git a/backend/internal/service/openai_ws_client.go b/backend/internal/service/openai_ws_client.go index d30c4a1cb3..c9b92d22e8 100644 --- a/backend/internal/service/openai_ws_client.go +++ b/backend/internal/service/openai_ws_client.go @@ -40,6 +40,13 @@ type openAIWSClientConn interface { Close() error } +// openAIWSIdlePingCapable is intentionally separate from openAIWSClientConn. +// A pool probe happens while no goroutine is reading an idle connection, which +// is not safe for every WebSocket implementation. +type openAIWSIdlePingCapable interface { + SupportsIdlePingWithoutReader() bool +} + // openAIWSClientDialer 抽象 WS 建连器。 type openAIWSClientDialer interface { Dial(ctx context.Context, wsURL string, headers http.Header, proxyURL string) (openAIWSClientConn, int, http.Header, error) @@ -329,6 +336,14 @@ func (c *coderOpenAIWSClientConn) Ping(ctx context.Context) error { return c.conn.Ping(ctx) } +// SupportsIdlePingWithoutReader reports the actual coder/websocket contract. +// Conn.Ping waits for a pong, while control frames are only consumed by Read. +// The pool deliberately has no reader on an idle connection, so using Ping as +// a health probe would deterministically time out a healthy socket. +func (*coderOpenAIWSClientConn) SupportsIdlePingWithoutReader() bool { + return false +} + func (c *coderOpenAIWSClientConn) Close() error { if c == nil || c.conn == nil { return nil diff --git a/backend/internal/service/openai_ws_client_test.go b/backend/internal/service/openai_ws_client_test.go index a88d626651..95614cdbfe 100644 --- a/backend/internal/service/openai_ws_client_test.go +++ b/backend/internal/service/openai_ws_client_test.go @@ -110,3 +110,7 @@ func TestCoderOpenAIWSClientDialer_ProxyTransportTLSHandshakeTimeout(t *testing. require.NotNil(t, transport) require.Equal(t, 10*time.Second, transport.TLSHandshakeTimeout) } + +func TestCoderOpenAIWSClientConn_DoesNotSupportIdlePingWithoutReader(t *testing.T) { + require.False(t, (&coderOpenAIWSClientConn{}).SupportsIdlePingWithoutReader()) +} diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index f7a54fb2b5..635312738b 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -18,6 +18,13 @@ import ( "github.com/tidwall/sjson" ) +func (s *OpenAIGatewayService) openAIWSIngressInterTurnIdleTimeout() time.Duration { + if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds <= 0 { + return 0 + } + return time.Duration(s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) * time.Second +} + func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( ctx context.Context, c *gin.Context, @@ -231,7 +238,11 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( if isCodexCLI { codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() } - codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) + codexBridgeEnabled := isCodexCLI && + !isOpenAIResponsesLiteWebSocketPayload(normalized) && + imageGenerationAllowed && + codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && + s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) if codexBridgeEnabled { payloadMap := make(map[string]any) if err := json.Unmarshal(normalized, &payloadMap); err != nil { @@ -360,8 +371,23 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } readClientMessage := func() ([]byte, error) { - msgType, payload, readErr := clientConn.Read(ctx) + readCtx := ctx + idleTimeout := s.openAIWSIngressInterTurnIdleTimeout() + cancelRead := func() {} + if idleTimeout > 0 { + readCtx, cancelRead = context.WithTimeout(ctx, idleTimeout) + } + msgType, payload, readErr := clientConn.Read(readCtx) + cancelRead() if readErr != nil { + if idleTimeout > 0 && errors.Is(readErr, context.DeadlineExceeded) && ctx.Err() == nil { + logOpenAIWSModeInfo("ingress_ws_inter_turn_idle_timeout account_id=%d timeout_seconds=%d", account.ID, int(idleTimeout.Seconds())) + return nil, NewOpenAIWSClientCloseError( + coderws.StatusNormalClosure, + "websocket idle timeout", + readErr, + ) + } return nil, readErr } if msgType != coderws.MessageText && msgType != coderws.MessageBinary { @@ -421,6 +447,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( storeDisabled, ) currentBridgePayload := firstPayload + // Keep the first turn as the stable conversation seed. The mapped model + // is resolved again for each turn below so an in-connection model switch + // cannot reuse another model's upstream cache identity. + grokCacheSeedPayload := firstPayload.payloadRaw var bridgeReplayInput []json.RawMessage bridgeReplayInputExists := false for turn := 1; ; turn++ { @@ -469,6 +499,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw), ) } + grokCacheIdentity := "" + if account.Platform == PlatformGrok { + grokCacheIdentity, err = resolveGrokWSCacheIdentity(c, account, grokCacheSeedPayload, currentBridgePayload.originalModel) + if err != nil { + return fmt.Errorf("resolve Grok websocket cache identity: %w", err) + } + } result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn( ctx, c, @@ -480,6 +517,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( currentBridgePayload.imageBillingModel, currentBridgePayload.imageSizeTier, currentBridgePayload.imageInputSize, + grokCacheIdentity, turn, writeClientMessage, ) @@ -1319,7 +1357,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( unpinSessionConn(sessionConnID) } } - shouldPreflightPing := turn > 1 && sessionLease != nil && turnRetry == 0 + shouldPreflightPing := turn > 1 && sessionLease != nil && sessionLease.SupportsIdlePingWithoutReader() && turnRetry == 0 if shouldPreflightPing && openAIWSIngressPreflightPingIdle > 0 && !lastTurnFinishedAt.IsZero() { if time.Since(lastTurnFinishedAt) < openAIWSIngressPreflightPingIdle { shouldPreflightPing = false diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go index 9ae1b855ea..5ec4b5333b 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go @@ -164,6 +164,111 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create") } +func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_IdleTimeoutReleasesStoreDisabledSession(t *testing.T) { + gin.SetMode(gin.TestMode) + + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + captureConn := &openAIWSCaptureConn{events: [][]byte{ + []byte(`{"type":"response.completed","response":{"id":"resp_idle_timeout","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`), + }} + captureDialer := &openAIWSCaptureDialer{conn: captureConn} + pool := newOpenAIWSConnPool(cfg) + pool.setClientDialerForTest(captureDialer) + defer pool.Close() + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: &httpUpstreamRecorder{}, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + openaiWSPool: pool, + } + account := &Account{ + ID: 116, + Name: "openai-ingress-idle-timeout", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{"responses_websockets_v2_enabled": true}, + } + + serverErrCh := make(chan error, 1) + wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.CloseNow() }() + + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + _, firstMessage, err := conn.Read(readCtx) + cancelRead() + if err != nil { + serverErrCh <- err + return + } + rec := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(rec) + ginCtx.Request = r.Clone(r.Context()) + serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) + })) + defer wsServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) + cancelDial() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, event, err := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + + select { + case proxyErr := <-serverErrCh: + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, proxyErr, &closeErr) + require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode()) + require.Equal(t, "websocket idle timeout", closeErr.Reason()) + case <-time.After(4 * time.Second): + t.Fatal("timed out waiting for idle ingress session to close") + } + + ap, ok := pool.getAccountPool(account.ID) + require.True(t, ok) + ap.mu.Lock() + require.Empty(t, ap.pinnedConns, "idle close must unpin a store=false session") + for _, conn := range ap.conns { + require.False(t, conn.isLeased(), "idle close must release the upstream lease") + } + ap.mu.Unlock() +} + func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCanOmitModel(t *testing.T) { gin.SetMode(gin.TestMode) @@ -298,7 +403,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCa require.Equal(t, "resp_omit_model_1", gjson.Get(requestToJSONString(captureConn.writes[1]), "previous_response_id").String()) } -func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImageBridge(t *testing.T) { +func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_CodexImageBridgeRespectsResponsesLite(t *testing.T) { gin.SetMode(gin.TestMode) cfg := &config.Config{} @@ -319,6 +424,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImag captureConn := &openAIWSCaptureConn{ events: [][]byte{ []byte(`{"type":"response.completed","response":{"id":"resp_codex_image_bridge","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`), + []byte(`{"type":"response.completed","response":{"id":"resp_codex_image_lite","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`), }, } captureDialer := &openAIWSCaptureDialer{conn: captureConn} @@ -418,6 +524,28 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImag require.Equal(t, coderws.MessageText, msgType) require.Equal(t, "resp_codex_image_bridge", gjson.GetBytes(message, "response.id").String()) + writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{ + "type":"response.create", + "model":"gpt-5.5", + "stream":false, + "previous_response_id":"resp_codex_image_bridge", + "client_metadata":{"ws_request_header_x_openai_internal_codex_responses_lite":"true"}, + "input":[ + {"type":"additional_tools","role":"developer","tools":[{"type":"custom","name":"exec","description":"Execute code-mode tools, including image_gen.imagegen."}]}, + {"type":"message","role":"user","content":[{"type":"input_text","text":"draw a cat"}]} + ] + }`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second) + msgType, message, err = clientConn.Read(readCtx) + cancelRead() + require.NoError(t, err) + require.Equal(t, coderws.MessageText, msgType) + require.Equal(t, "resp_codex_image_lite", gjson.GetBytes(message, "response.id").String()) + _ = clientConn.Close(coderws.StatusNormalClosure, "done") select { @@ -427,12 +555,19 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImag t.Fatal("等待 ingress websocket 结束超时") } - require.Len(t, captureConn.writes, 1) - upstreamPayload := requestToJSONString(captureConn.writes[0]) - require.True(t, gjson.Get(upstreamPayload, `tools.#(type=="image_generation")`).Exists()) - require.Equal(t, "png", gjson.Get(upstreamPayload, `tools.#(type=="image_generation").output_format`).String()) - require.Equal(t, "auto", gjson.Get(upstreamPayload, "tool_choice").String()) - require.Contains(t, gjson.Get(upstreamPayload, "instructions").String(), "image_generation") + require.Len(t, captureConn.writes, 2) + nonLitePayload := requestToJSONString(captureConn.writes[0]) + require.True(t, gjson.Get(nonLitePayload, `tools.#(type=="image_generation")`).Exists()) + require.Equal(t, "png", gjson.Get(nonLitePayload, `tools.#(type=="image_generation").output_format`).String()) + require.Equal(t, "auto", gjson.Get(nonLitePayload, "tool_choice").String()) + require.Contains(t, gjson.Get(nonLitePayload, "instructions").String(), "image_generation") + + litePayload := requestToJSONString(captureConn.writes[1]) + require.False(t, gjson.Get(litePayload, `tools.#(type=="image_generation")`).Exists()) + require.False(t, gjson.Get(litePayload, "tool_choice").Exists()) + require.NotContains(t, gjson.Get(litePayload, "instructions").String(), "image_generation") + require.Equal(t, "exec", gjson.Get(litePayload, `input.#(type=="additional_tools").tools.0.name`).String()) + require.Contains(t, gjson.Get(litePayload, `input.#(type=="additional_tools").tools.0.description`).String(), "image_gen.imagegen") } func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_DedicatedModeDoesNotReuseConnAcrossSessions(t *testing.T) { diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index 67fdddc9f7..511b28d0ba 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -86,6 +86,11 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( if v := strings.TrimSpace(c.Request.Header.Get("accept-language")); v != "" { headers.Set("accept-language", v) } + for _, value := range c.Request.Header.Values("x-codex-beta-features") { + if value = strings.TrimSpace(value); value != "" { + headers.Add("x-codex-beta-features", value) + } + } } // OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。 if account != nil && account.Type == AccountTypeOAuth { diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index adae109e09..bb4ac2242c 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -602,6 +602,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T c.Request.Header.Set("User-Agent", "codex_cli_rs/0.98.0") c.Request.Header.Set("session_id", "sess-oauth-1") c.Request.Header.Set("conversation_id", "conv-oauth-1") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") cfg := &config.Config{} cfg.Security.URLAllowlist.Enabled = false @@ -661,6 +662,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T require.True(t, gjson.Get(requestJSON, "stream").Exists(), "WSv2 payload 应保留 stream 字段") require.True(t, gjson.Get(requestJSON, "stream").Bool(), "OAuth Codex 规范化后应强制 stream=true") require.Equal(t, openAIWSBetaV2Value, captureDialer.lastHeaders.Get("OpenAI-Beta")) + require.Equal(t, "remote_compaction_v2", captureDialer.lastHeaders.Get("x-codex-beta-features")) // OAuth 账号的 session_id/conversation_id 应被 isolateOpenAISessionID 隔离, // 测试中未设置 api_key 到 context,apiKeyID=0。 require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id")) diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index dce0c7b6db..dbc24c850a 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -12,6 +12,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" ) @@ -155,6 +156,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( imageBillingModel string, imageSizeTier string, imageInputSize string, + grokCacheIdentity string, turn int, writeClientMessage func([]byte) error, ) (*OpenAIForwardResult, error) { @@ -179,21 +181,19 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) var upstreamReq *http.Request if account.Platform == PlatformGrok { - upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) - if originalModel != "" { - if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" { - upstreamModel = mappedModel - } - } - if upstreamModel == "" { - upstreamModel = "grok-4.3" - } + upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel) + grokIntentSourceBody := body body, err = patchGrokResponsesBody(body, upstreamModel) if err != nil { releaseUpstreamCtx() return nil, err } - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token) + body, err = applyGrokResponsesCacheIdentity(body, grokIntentSourceBody, grokCacheIdentity, account.IsGrokOAuth()) + if err != nil { + releaseUpstreamCtx() + return nil, fmt.Errorf("apply grok prompt cache identity: %w", err) + } + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity) } else { upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) } @@ -201,6 +201,9 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( if err != nil { return nil, err } + if account.Platform != PlatformGrok && isOpenAIResponsesLiteWebSocketPayload(payload) { + upstreamReq.Header.Set(responsesLiteHeader, "true") + } proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -222,6 +225,9 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( if resp.StatusCode >= 400 { respBody, _ := io.ReadAll(io.LimitReader(resp.Body, openAIWSHTTPBridgeErrorBodyLimitBytes)) + if account.Platform == PlatformGrok { + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + } upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) if upstreamMsg == "" { upstreamMsg = http.StatusText(resp.StatusCode) @@ -229,6 +235,9 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( _ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(resp.StatusCode, upstreamMsg)) return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg) } + if account.Platform == PlatformGrok { + s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + } responseID := "" usage := OpenAIUsage{} @@ -407,3 +416,25 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( } return resultWithUsage(), errors.New("upstream http bridge stream ended before terminal event") } + +func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, payload []byte, originalModel string) (string, error) { + body, err := prepareOpenAIWSHTTPBridgeBody(payload) + if err != nil { + return "", err + } + upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel) + return resolveGrokCacheIdentity(c, body, "", upstreamModel), nil +} + +func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string { + upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if account != nil && originalModel != "" { + if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" { + upstreamModel = mappedModel + } + } + if upstreamModel == "" { + upstreamModel = grokDefaultResponsesModel + } + return upstreamModel +} diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 4d1e4a0374..b1bfb4d342 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -3,6 +3,7 @@ package service import ( "context" "errors" + "fmt" "io" "net/http" "net/http/httptest" @@ -90,7 +91,7 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) { Concurrency: 1, Status: StatusActive, } - payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"hi"}`) + payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"client_metadata":{"ws_request_header_x_openai_internal_codex_responses_lite":"true"},"input":"hi"}`) type bridgeResult struct { result *OpenAIForwardResult @@ -127,6 +128,7 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) { "", "", "", + "", 1, writeClient, ) @@ -171,29 +173,82 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) { require.NotNil(t, upstream.lastReq) require.Equal(t, http.MethodPost, upstream.lastReq.Method) + require.Equal(t, "true", upstream.lastReq.Header.Get(responsesLiteHeader)) require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) } +func TestProxyOpenAIWSHTTPBridgeTurnForGrokDefaultsEmptyModelTo45(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_grok_default","model":"grok-4.5"}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_grok_default","model":"grok-4.5","usage":{"input_tokens":1,"output_tokens":1}}}`, + "", + }, "\n"))), + }} + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, + httpUpstream: upstream, + } + account := &Account{ + ID: 72, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL}, + } + payload := []byte(`{"type":"response.create","generate":true,"stream":true,"input":"hi"}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + var events [][]byte + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "access-token", payload, len(payload), + "", "", "", "", "", 1, + func(message []byte) error { + events = append(events, append([]byte(nil), message...)) + return nil + }, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String()) + require.Len(t, events, 2) +} + func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) { gin.SetMode(gin.TestMode) - sseBody := strings.Join([]string{ - `data: {"type":"response.created","response":{"id":"resp_grok_ws","model":"grok-4.3"}}`, - "", - `data: {"type":"response.output_text.delta","response":{"id":"resp_grok_ws"},"delta":"ok"}`, - "", - `data: {"type":"response.completed","response":{"id":"resp_grok_ws","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2}}}`, - "", - }, "\n") - upstream := &httpUpstreamRecorder{resp: &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{ - "Content-Type": []string{"text/event-stream"}, - "Xai-Request-Id": []string{"xai-ws-req"}, - }, - Body: io.NopCloser(strings.NewReader(sseBody)), + bridgeResponse := func(responseID, requestID string, cachedTokens int) *http.Response { + sseBody := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"` + responseID + `","model":"grok-4.3"}}`, + "", + `data: {"type":"response.output_text.delta","response":{"id":"` + responseID + `"},"delta":"ok"}`, + "", + `data: {"type":"response.completed","response":{"id":"` + responseID + `","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2,"input_tokens_details":{"cached_tokens":` + fmt.Sprintf("%d", cachedTokens) + `}}}}`, + "", + }, "\n") + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "Xai-Request-Id": []string{requestID}, + }, + Body: io.NopCloser(strings.NewReader(sseBody)), + } + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + bridgeResponse("resp_grok_ws_1", "xai-ws-req-1", 0), + bridgeResponse("resp_grok_ws_2", "xai-ws-req-2", 3), + bridgeResponse("resp_grok_ws_3", "xai-ws-req-3", 0), }} svc := &OpenAIGatewayService{ cfg: &config.Config{ @@ -241,6 +296,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) req := r.Clone(r.Context()) req.Header = req.Header.Clone() ginCtx.Request = req + ginCtx.Set("api_key", &APIKey{ID: 7101}) errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, nil) })) @@ -271,6 +327,33 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String()) require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) + require.Equal(t, "resp_grok_ws_1", gjson.GetBytes(completed, "response.id").String()) + + writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`)) + cancelWrite() + require.NoError(t, err) + + created = readEvent() + delta = readEvent() + completed = readEvent() + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String()) + require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) + require.Equal(t, "resp_grok_ws_2", gjson.GetBytes(completed, "response.id").String()) + + writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok-4.3","stream":true,"previous_response_id":"resp_grok_ws_2","input":"third turn with a different model"}`)) + cancelWrite() + require.NoError(t, err) + + created = readEvent() + delta = readEvent() + completed = readEvent() + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String()) + require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) + require.Equal(t, "resp_grok_ws_3", gjson.GetBytes(completed, "response.id").String()) _ = clientConn.Close(coderws.StatusNormalClosure, "done") select { @@ -280,10 +363,30 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) require.Fail(t, "proxy did not finish after client close") } + require.Len(t, upstream.requests, 3) + require.Len(t, upstream.bodies, 3) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) - require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String()) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[1], "model").String()) + require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[2], "model").String()) + require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()) + require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader)) + require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) + require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String()) + firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String() + secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String() + thirdIdentity := gjson.GetBytes(upstream.bodies[2], "prompt_cache_key").String() + require.NotEmpty(t, firstIdentity) + require.Equal(t, firstIdentity, secondIdentity) + require.NotEmpty(t, thirdIdentity) + require.NotEqual(t, firstIdentity, thirdIdentity) + require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader)) + require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader)) + require.Equal(t, thirdIdentity, upstream.requests[2].Header.Get(grokConversationIDHeader)) require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists()) @@ -575,3 +678,100 @@ func TestOpenAIWSHTTPBridgeKeepsContinuationFramesOnHTTPWithoutPreviousResponseI require.Equal(t, 0, captureDialer.DialCount()) require.Empty(t, captureConn.writes) } + +func TestOpenAIWSHTTPBridge_IdleTimeoutClosesClientSession(t *testing.T) { + gin.SetMode(gin.TestMode) + + sseBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_bridge_idle","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`, + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(sseBody)), + }} + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true + cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 + cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 + + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: upstream, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + } + account := &Account{ + ID: 20, + Name: "api-key-bridge-idle-timeout", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-upstream"}, + Extra: map[string]any{"responses_websockets_v2_enabled": true}, + Concurrency: 1, + Status: StatusActive, + Schedulable: true, + } + + errCh := make(chan error, 1) + wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) + if err != nil { + errCh <- err + return + } + defer func() { _ = conn.CloseNow() }() + + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + _, firstMessage, err := conn.Read(readCtx) + cancelRead() + if err != nil { + errCh <- err + return + } + rec := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(rec) + ginCtx.Request = r.Clone(r.Context()) + errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) + })) + defer wsServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) + cancelDial() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"input":"hello"}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, event, err := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + + select { + case proxyErr := <-errCh: + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, proxyErr, &closeErr) + require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode()) + require.Equal(t, "websocket idle timeout", closeErr.Reason()) + case <-time.After(4 * time.Second): + t.Fatal("timed out waiting for idle HTTP bridge session to close") + } + require.Len(t, upstream.bodies, 1, "an idle client must not leave a continuation request running") +} diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 3a02da6fe5..8579cc58d6 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -208,6 +208,14 @@ func (l *openAIWSConnLease) PingWithTimeout(timeout time.Duration) error { return conn.pingWithTimeout(timeout) } +func (l *openAIWSConnLease) SupportsIdlePingWithoutReader() bool { + conn, err := l.activeConn() + if err != nil { + return false + } + return conn.supportsIdlePingWithoutReader() +} + func (l *openAIWSConnLease) MarkBroken() { if l == nil || l.pool == nil || l.conn == nil || l.released.Load() { return @@ -223,6 +231,9 @@ func (l *openAIWSConnLease) Release() { return } l.conn.release() + if l.pool != nil { + l.pool.notifyAccountPoolChanged(l.accountID) + } } type openAIWSConn struct { @@ -230,6 +241,7 @@ type openAIWSConn struct { ws openAIWSClientConn handshakeHeaders http.Header + betaFeatures string leaseCh chan struct{} closedCh chan struct{} @@ -438,6 +450,16 @@ func (c *openAIWSConn) pingWithTimeout(timeout time.Duration) error { return nil } +func (c *openAIWSConn) supportsIdlePingWithoutReader() bool { + if c == nil || c.ws == nil { + return false + } + capable, ok := c.ws.(openAIWSIdlePingCapable) + // Test and alternate implementations keep the historical probe behavior + // unless they explicitly declare it unsafe. + return !ok || capable.SupportsIdlePingWithoutReader() +} + func (c *openAIWSConn) touch() { if c == nil { return @@ -503,6 +525,10 @@ func (c *openAIWSConn) handshakeHeader(name string) string { return strings.TrimSpace(c.handshakeHeaders.Get(strings.TrimSpace(name))) } +func (c *openAIWSConn) matchesBetaFeatures(betaFeatures string) bool { + return c != nil && c.betaFeatures == betaFeatures +} + func (c *openAIWSConn) isPrewarmed() bool { if c == nil { return false @@ -521,6 +547,7 @@ type openAIWSAccountPool struct { mu sync.Mutex conns map[string]*openAIWSConn pinnedConns map[string]int + changedCh chan struct{} creating int generation uint64 lastCleanupAt time.Time @@ -531,6 +558,23 @@ type openAIWSAccountPool struct { prewarmFailAt time.Time } +func (ap *openAIWSAccountPool) changeChannelLocked() chan struct{} { + if ap.changedCh == nil { + ap.changedCh = make(chan struct{}) + } + return ap.changedCh +} + +func (ap *openAIWSAccountPool) signalChangedLocked() { + if ap == nil { + return + } + if ap.changedCh != nil { + close(ap.changedCh) + } + ap.changedCh = make(chan struct{}) +} + type OpenAIWSPoolMetricsSnapshot struct { AcquireTotal int64 AcquireReuseTotal int64 @@ -687,7 +731,7 @@ func (p *openAIWSConnPool) runBackgroundPingSweep() { g.SetLimit(10) for _, item := range candidates { item := item - if item.conn == nil || item.conn.isLeased() || item.conn.waiters.Load() > 0 { + if item.conn == nil || item.conn.isLeased() || item.conn.waiters.Load() > 0 || !item.conn.supportsIdlePingWithoutReader() { continue } g.Go(func() error { @@ -792,7 +836,9 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errors.New("ws url is empty") } +retryAcquire: accountID := req.Account.ID + betaFeatures := normalizeOpenAIWSBetaFeatures(req.Headers) effectiveMaxConns := p.effectiveMaxConnsByAccount(req.Account) if effectiveMaxConns <= 0 { return nil, errOpenAIWSConnQueueFull @@ -820,7 +866,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errOpenAIWSPreferredConnUnavailable } preferredConn, ok := ap.conns[preferredConnID] - if !ok || preferredConn == nil { + if !ok || !preferredConn.matchesBetaFeatures(betaFeatures) { p.recordConnPickDuration(time.Since(pickStartedAt)) ap.mu.Unlock() closeOpenAIWSConns(evicted) @@ -901,7 +947,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok && conn.tryAcquire() { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) && conn.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) ap.mu.Unlock() @@ -923,7 +969,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } } - best := p.pickLeastBusyConnLocked(ap, "") + best := p.pickLeastBusyConnLocked(ap, "", betaFeatures) if best != nil && best.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) @@ -945,7 +991,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return lease, nil } for _, conn := range ap.conns { - if conn == nil || conn == best { + if conn == nil || conn == best || !conn.matchesBetaFeatures(betaFeatures) { continue } if conn.tryAcquire() { @@ -971,6 +1017,37 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } } + if !req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns { + compatible := p.pickLeastBusyConnLocked(ap, "", betaFeatures) + if idle := p.pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap, betaFeatures); idle != nil { + delete(ap.conns, idle.id) + evicted = append(evicted, idle) + p.metrics.scaleDownTotal.Add(1) + } else if compatible == nil { + hasConnection := false + for _, conn := range ap.conns { + if conn != nil { + hasConnection = true + break + } + } + if !hasConnection && ap.creating == 0 { + ap.mu.Unlock() + closeOpenAIWSConns(evicted) + return nil, errOpenAIWSConnClosed + } + changedCh := ap.changeChannelLocked() + ap.mu.Unlock() + closeOpenAIWSConns(evicted) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-changedCh: + goto retryAcquire + } + } + } + if req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns { if idle := p.pickOldestIdleConnLocked(ap); idle != nil { delete(ap.conns, idle.id) @@ -994,6 +1071,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque if dialErr != nil { ap.prewarmFails++ ap.prewarmFailAt = time.Now() + ap.signalChangedLocked() ap.mu.Unlock() return nil, dialErr } @@ -1022,7 +1100,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errOpenAIWSConnQueueFull } - target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID) + target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, betaFeatures) connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) if target == nil { @@ -1095,6 +1173,22 @@ func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *op return oldest } +func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap *openAIWSAccountPool, betaFeatures string) *openAIWSConn { + if ap == nil || len(ap.conns) == 0 { + return nil + } + var oldest *openAIWSConn + for _, conn := range ap.conns { + if conn == nil || conn.matchesBetaFeatures(betaFeatures) || conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) { + continue + } + if oldest == nil || conn.lastUsedAt().Before(oldest.lastUsedAt()) { + oldest = conn + } + } + return oldest +} + func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAccountPool { if p == nil || accountID <= 0 { return nil @@ -1107,6 +1201,7 @@ func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAcco ap := &openAIWSAccountPool{ conns: make(map[string]*openAIWSConn), pinnedConns: make(map[string]int), + changedCh: make(chan struct{}), } actual, _ := p.accounts.LoadOrStore(accountID, ap) if typed, ok := actual.(*openAIWSAccountPool); ok && typed != nil { @@ -1132,6 +1227,16 @@ func (p *openAIWSConnPool) getAccountPool(accountID int64) (*openAIWSAccountPool return ap, typed && ap != nil } +func (p *openAIWSConnPool) notifyAccountPoolChanged(accountID int64) { + ap, ok := p.getAccountPool(accountID) + if !ok || ap == nil { + return + } + ap.mu.Lock() + ap.signalChangedLocked() + ap.mu.Unlock() +} + func (p *openAIWSConnPool) isConnPinnedLocked(ap *openAIWSAccountPool, connID string) bool { if ap == nil || connID == "" || len(ap.pinnedConns) == 0 { return false @@ -1218,17 +1323,20 @@ func (p *openAIWSConnPool) cleanupAccountLocked(ap *openAIWSAccountPool, now tim p.metrics.scaleDownTotal.Add(int64(redundant)) } } + if len(evicted) > 0 { + ap.signalChangedLocked() + } return evicted } -func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID string) *openAIWSConn { +func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID, betaFeatures string) *openAIWSConn { if ap == nil || len(ap.conns) == 0 { return nil } preferredConnID = stringsTrim(preferredConnID) if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) { return conn } } @@ -1236,7 +1344,7 @@ func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, pref var bestWaiters int32 var bestLastUsed time.Time for _, conn := range ap.conns { - if conn == nil { + if conn == nil || !conn.matchesBetaFeatures(betaFeatures) { continue } waiters := conn.waiters.Load() @@ -1407,6 +1515,7 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ if err != nil { ap.prewarmFails++ ap.prewarmFailAt = time.Now() + ap.signalChangedLocked() ap.mu.Unlock() continue } @@ -1416,6 +1525,7 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ continue } if len(ap.conns) >= p.effectiveMaxConnsByAccount(req.Account) { + ap.signalChangedLocked() ap.mu.Unlock() conn.close() continue @@ -1423,6 +1533,7 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ ap.conns[conn.id] = conn ap.prewarmFails = 0 ap.prewarmFailAt = time.Time{} + ap.signalChangedLocked() ap.mu.Unlock() } } @@ -1470,6 +1581,7 @@ func (p *openAIWSConnPool) evictConn(accountID int64, connID string) { if len(ap.pinnedConns) > 0 { delete(ap.pinnedConns, connID) } + ap.signalChangedLocked() } ap.mu.Unlock() } @@ -1522,9 +1634,11 @@ func (p *openAIWSConnPool) UnpinConn(accountID int64, connID string) { count := ap.pinnedConns[connID] if count <= 1 { delete(ap.pinnedConns, connID) + ap.signalChangedLocked() return } ap.pinnedConns[connID] = count - 1 + ap.signalChangedLocked() } func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequest) (*openAIWSConn, error) { @@ -1561,7 +1675,9 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ } } id := p.nextConnID(req.Account.ID) - return newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders), nil + pooledConn := newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders) + pooledConn.betaFeatures = normalizeOpenAIWSBetaFeatures(req.Headers) + return pooledConn, nil } func (p *openAIWSConnPool) nextConnID(accountID int64) string { @@ -1575,7 +1691,7 @@ func (p *openAIWSConnPool) nextConnID(accountID int64) string { } func (p *openAIWSConnPool) shouldHealthCheckConn(conn *openAIWSConn) bool { - if conn == nil { + if conn == nil || !conn.supportsIdlePingWithoutReader() { return false } return conn.idleDuration(time.Now()) >= openAIWSConnHealthCheckIdle @@ -1631,7 +1747,7 @@ func (p *openAIWSConnPool) effectiveMaxConnsByAccount(account *Account) int { if account.Concurrency <= 0 { return 0 } - return account.Concurrency + return min(account.Concurrency, hardCap) } if account == nil || !p.dynamicMaxConnsEnabled() { return hardCap @@ -1739,6 +1855,31 @@ func cloneOpenAIWSAcquireRequestPtr(req *openAIWSAcquireRequest) *openAIWSAcquir return &copied } +func normalizeOpenAIWSBetaFeatures(headers http.Header) string { + features := make(map[string]struct{}) + for name, values := range headers { + if !strings.EqualFold(strings.TrimSpace(name), "x-codex-beta-features") { + continue + } + for _, value := range values { + for _, feature := range strings.Split(value, ",") { + if feature = strings.TrimSpace(feature); feature != "" { + features[feature] = struct{}{} + } + } + } + } + if len(features) == 0 { + return "" + } + normalized := make([]string, 0, len(features)) + for feature := range features { + normalized = append(normalized, feature) + } + sort.Strings(normalized) + return strings.Join(normalized, ",") +} + func cloneHeader(src http.Header) http.Header { if src == nil { return nil diff --git a/backend/internal/service/openai_ws_pool_test.go b/backend/internal/service/openai_ws_pool_test.go index b2683ee041..8d339359ee 100644 --- a/backend/internal/service/openai_ws_pool_test.go +++ b/backend/internal/service/openai_ws_pool_test.go @@ -342,6 +342,171 @@ func TestOpenAIWSConnPool_ForceNewConnSkipsReuse(t *testing.T) { require.Equal(t, 2, dialer.DialCount(), "ForceNewConn=true 时应跳过空闲连接复用并新建连接") } +func TestOpenAIWSConnPool_AcquireReusesOnlyMatchingBetaFeatures(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + + account := &Account{ID: 128, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + } + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {" remote_compaction_v2 ", " responses_websockets_v2 "}} + betaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + require.False(t, betaLease.Reused()) + require.NotEqual(t, plainConnID, betaLease.ConnID()) + betaConnID := betaLease.ConnID() + betaLease.Release() + + reorderedReq := baseReq + reorderedReq.Headers = http.Header{"X-Codex-Beta-Features": {"responses_websockets_v2,remote_compaction_v2"}} + reorderedLease, err := pool.Acquire(context.Background(), reorderedReq) + require.NoError(t, err) + require.True(t, reorderedLease.Reused()) + require.Equal(t, betaConnID, reorderedLease.ConnID()) + reorderedLease.Release() + + _, err = pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: baseReq.WSURL, + Headers: betaReq.Headers, + PreferredConnID: plainConnID, + ForcePreferredConn: true, + }) + require.ErrorIs(t, err, errOpenAIWSPreferredConnUnavailable) + + plainLease, err = pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + require.True(t, plainLease.Reused()) + require.Equal(t, plainConnID, plainLease.ConnID()) + plainLease.Release() + + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireReplacesIdleConnWithDifferentBetaFeatures(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + + account := &Account{ID: 129, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + plainLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + }) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}}, + }) + require.NoError(t, err) + require.False(t, betaLease.Reused()) + require.NotEqual(t, plainConnID, betaLease.ConnID()) + betaLease.Release() + + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireWaitsForBusyIncompatibleConnection(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 130, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + + type acquireResult struct { + lease *openAIWSConnLease + err error + } + resultCh := make(chan acquireResult, 1) + var done atomic.Bool + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + go func() { + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}} + lease, acquireErr := pool.Acquire(ctx, betaReq) + resultCh <- acquireResult{lease: lease, err: acquireErr} + done.Store(true) + }() + + require.Never(t, done.Load, 50*time.Millisecond, 5*time.Millisecond) + plainLease.Release() + + result := <-resultCh + require.NoError(t, result.err) + require.NotNil(t, result.lease) + require.False(t, result.lease.Reused()) + require.NotEqual(t, plainConnID, result.lease.ConnID()) + result.lease.Release() + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireReplacesIncompatibleIdleWhenMatchingBusy(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 131, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}} + busyBetaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + + secondBetaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + require.False(t, secondBetaLease.Reused()) + require.NotEqual(t, plainConnID, secondBetaLease.ConnID()) + require.NotEqual(t, busyBetaLease.ConnID(), secondBetaLease.ConnID()) + + secondBetaLease.Release() + busyBetaLease.Release() + require.Equal(t, 3, dialer.DialCount()) +} + func TestOpenAIWSConnPool_AcquireForcePreferredConnUnavailable(t *testing.T) { cfg := &config.Config{} cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 @@ -574,7 +739,7 @@ func TestOpenAIWSConnPool_EffectiveMaxConnsDisabledFallbackHardCap(t *testing.T) require.Equal(t, 8, pool.effectiveMaxConnsByAccount(account), "关闭动态模式后应保持旧行为") } -func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2UsesAccountConcurrency(t *testing.T) { +func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2RespectsHardCap(t *testing.T) { cfg := &config.Config{} cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 8 @@ -585,7 +750,7 @@ func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2UsesAccountConc pool := newOpenAIWSConnPool(cfg) high := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 20} - require.Equal(t, 20, pool.effectiveMaxConnsByAccount(high), "v2 路径应直接使用账号并发数作为池上限") + require.Equal(t, 8, pool.effectiveMaxConnsByAccount(high), "v2 路径也必须受连接池硬上限约束") nonPositive := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 0} require.Equal(t, 0, pool.effectiveMaxConnsByAccount(nonPositive), "并发数<=0 时应不可调度") @@ -1154,6 +1319,9 @@ func TestOpenAIWSConnPool_UtilityBranches(t *testing.T) { conn := newOpenAIWSConn("health", 1, &openAIWSFakeConn{}, nil) conn.lastUsedNano.Store(time.Now().Add(-openAIWSConnHealthCheckIdle - time.Second).UnixNano()) require.True(t, pool.shouldHealthCheckConn(conn)) + unsafeConn := newOpenAIWSConn("unsafe_health", 1, &openAIWSIdlePingUnsupportedConn{}, nil) + unsafeConn.lastUsedNano.Store(time.Now().Add(-openAIWSConnHealthCheckIdle - time.Second).UnixNano()) + require.False(t, pool.shouldHealthCheckConn(unsafeConn)) } func TestOpenAIWSConn_LeaseAndTimeHelpers_NilAndClosedBranches(t *testing.T) { @@ -1445,6 +1613,14 @@ type openAIWSPingBlockingConn struct { release <-chan struct{} } +type openAIWSIdlePingUnsupportedConn struct { + openAIWSFakeConn +} + +func (c *openAIWSIdlePingUnsupportedConn) SupportsIdlePingWithoutReader() bool { + return false +} + func (c *openAIWSPingBlockingConn) WriteJSON(context.Context, any) error { return nil } diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go index e33dcf82a8..d95bbadbb8 100644 --- a/backend/internal/service/ops_models.go +++ b/backend/internal/service/ops_models.go @@ -8,6 +8,7 @@ import ( type OpsSystemLog struct { ID int64 `json:"id"` CreatedAt time.Time `json:"created_at"` + Host string `json:"host"` Level string `json:"level"` Component string `json:"component"` Message string `json:"message"` diff --git a/backend/internal/service/ops_port.go b/backend/internal/service/ops_port.go index 46d171c7c3..2b73d4a694 100644 --- a/backend/internal/service/ops_port.go +++ b/backend/internal/service/ops_port.go @@ -194,6 +194,7 @@ type OpsInsertSystemMetricsInput struct { type OpsInsertSystemLogInput struct { CreatedAt time.Time + Host string Level string Component string Message string @@ -210,6 +211,7 @@ type OpsInsertSystemLogInput struct { type OpsSystemLogFilter struct { StartTime *time.Time EndTime *time.Time + Host string Level string Component string @@ -230,6 +232,7 @@ type OpsSystemLogFilter struct { type OpsSystemLogCleanupFilter struct { StartTime *time.Time EndTime *time.Time + Host string Level string Component string diff --git a/backend/internal/service/ops_system_log_service.go b/backend/internal/service/ops_system_log_service.go index b3be37e8ae..b96ae89d92 100644 --- a/backend/internal/service/ops_system_log_service.go +++ b/backend/internal/service/ops_system_log_service.go @@ -89,6 +89,7 @@ func marshalSystemLogCleanupConditions(filter *OpsSystemLogCleanupFilter) string return "{}" } payload := map[string]any{ + "host": strings.TrimSpace(filter.Host), "level": strings.TrimSpace(filter.Level), "component": strings.TrimSpace(filter.Component), "request_id": strings.TrimSpace(filter.RequestID), diff --git a/backend/internal/service/ops_system_log_service_test.go b/backend/internal/service/ops_system_log_service_test.go index 8b5a84c1f0..e8c6199f17 100644 --- a/backend/internal/service/ops_system_log_service_test.go +++ b/backend/internal/service/ops_system_log_service_test.go @@ -101,6 +101,7 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) { now := time.Now().UTC() filter := &OpsSystemLogCleanupFilter{ StartTime: &now, + Host: "api-node-1", Level: "warn", RequestID: "req-1", ClientRequestID: "creq-1", @@ -119,6 +120,9 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) { if audit == nil { t.Fatalf("expected cleanup audit") } + if !strings.Contains(audit.Conditions, `"host":"api-node-1"`) { + t.Fatalf("audit conditions should include host: %s", audit.Conditions) + } if !strings.Contains(audit.Conditions, `"client_request_id":"creq-1"`) { t.Fatalf("audit conditions should include client_request_id: %s", audit.Conditions) } diff --git a/backend/internal/service/ops_system_log_sink.go b/backend/internal/service/ops_system_log_sink.go index 2ff273be53..2e6f5515c8 100644 --- a/backend/internal/service/ops_system_log_sink.go +++ b/backend/internal/service/ops_system_log_sink.go @@ -27,6 +27,7 @@ type OpsSystemLogSinkHealth struct { type OpsSystemLogSink struct { opsRepo OpsRepository + host string queue chan *logger.LogEvent @@ -45,10 +46,14 @@ type OpsSystemLogSink struct { lastError atomic.Value } +const maxSystemLogHostLength = 255 + func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink { ctx, cancel := context.WithCancel(context.Background()) + rawHost, err := os.Hostname() s := &OpsSystemLogSink{ opsRepo: opsRepo, + host: normalizeSystemLogHost(rawHost, err), queue: make(chan *logger.LogEvent, 5000), batchSize: 200, flushInterval: time.Second, @@ -59,6 +64,18 @@ func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink { return s } +func normalizeSystemLogHost(host string, err error) string { + host = strings.TrimSpace(host) + if err != nil || host == "" { + return "unknown" + } + runes := []rune(host) + if len(runes) > maxSystemLogHostLength { + return string(runes[:maxSystemLogHostLength]) + } + return host +} + func (s *OpsSystemLogSink) Start() { if s == nil || s.opsRepo == nil { return @@ -220,6 +237,7 @@ func (s *OpsSystemLogSink) flushBatch(baseCtx context.Context, batch []*logger.L inputs = append(inputs, &OpsInsertSystemLogInput{ CreatedAt: createdAt, + Host: s.host, Level: strings.ToLower(strings.TrimSpace(event.Level)), Component: component, Message: message, diff --git a/backend/internal/service/ops_system_log_sink_test.go b/backend/internal/service/ops_system_log_sink_test.go index b43d44c32e..0d15f1a662 100644 --- a/backend/internal/service/ops_system_log_sink_test.go +++ b/backend/internal/service/ops_system_log_sink_test.go @@ -140,6 +140,7 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) { } sink := NewOpsSystemLogSink(repo) + sink.host = "api-node-1" sink.batchSize = 1 sink.flushInterval = 10 * time.Millisecond sink.Start() @@ -172,6 +173,9 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) { t.Fatalf("captured len = %d, want 1", len(captured)) } item := captured[0] + if item.Host != "api-node-1" { + t.Fatalf("host = %q, want api-node-1", item.Host) + } if item.RequestID != "req-1" || item.ClientRequestID != "creq-1" { t.Fatalf("unexpected request ids: %+v", item) } @@ -324,3 +328,20 @@ func TestOpsSystemLogSink_HelperFunctions(t *testing.T) { } } } + +func TestNormalizeSystemLogHost(t *testing.T) { + if got := normalizeSystemLogHost(" api-node-1 ", nil); got != "api-node-1" { + t.Fatalf("trimmed host = %q, want api-node-1", got) + } + if got := normalizeSystemLogHost("", nil); got != "unknown" { + t.Fatalf("empty host = %q, want unknown", got) + } + if got := normalizeSystemLogHost("api-node-1", errors.New("hostname unavailable")); got != "unknown" { + t.Fatalf("errored host = %q, want unknown", got) + } + longHost := strings.Repeat("节", maxSystemLogHostLength+1) + got := normalizeSystemLogHost(longHost, nil) + if runeCount := len([]rune(got)); runeCount != maxSystemLogHostLength { + t.Fatalf("truncated host rune count = %d, want %d", runeCount, maxSystemLogHostLength) + } +} diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index 04feb8002a..9b7ee08990 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/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/shopspring/decimal" ) @@ -445,7 +446,9 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen IsMobile: req.IsMobile, ReturnURL: providerReturnURL, }, sel, outTradeNo, payAmountStr, subject) + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") pr, err := prov.CreatePayment(ctx, providerReq) + finishProviderCall() if err != nil { slog.Error("[PaymentService] CreatePayment failed", "provider", sel.ProviderKey, "instance", sel.InstanceID, "error", err) if appErr := new(infraerrors.ApplicationError); errors.As(err, &appErr) { diff --git a/backend/internal/service/payment_order_lifecycle.go b/backend/internal/service/payment_order_lifecycle.go index 8ed18797dd..46a2e00605 100644 --- a/backend/internal/service/payment_order_lifecycle.go +++ b/backend/internal/service/payment_order_lifecycle.go @@ -13,6 +13,7 @@ import ( "github.com/Wei-Shaw/sub2api/ent/paymentorder" "github.com/Wei-Shaw/sub2api/internal/payment" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // --- Cancel & Expire --- @@ -157,7 +158,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym if queryRef == "" { return "" } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.QueryOrder(ctx, queryRef) + finishProviderCall() if err != nil { slog.Warn("query upstream failed", "orderID", o.ID, "error", err) return "" @@ -199,7 +202,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym return "" } if cp, ok := prov.(payment.CancelableProvider); ok { + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") _ = cp.CancelPayment(ctx, queryRef) + finishProviderCall() } return "" } @@ -208,7 +213,9 @@ func requeryPaidOrderOnce(ctx context.Context, prov payment.Provider, queryRef s if prov == nil || strings.TrimSpace(queryRef) == "" { return nil, false } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.QueryOrder(ctx, queryRef) + finishProviderCall() if err != nil { slog.Warn("query upstream retry failed", "queryRef", queryRef, "error", err) return nil, false diff --git a/backend/internal/service/payment_refund.go b/backend/internal/service/payment_refund.go index 91822680ed..bc073a2c34 100644 --- a/backend/internal/service/payment_refund.go +++ b/backend/internal/service/payment_refund.go @@ -19,6 +19,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/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // --- Refund Flow --- @@ -347,12 +348,14 @@ func (s *PaymentService) gwRefund(ctx context.Context, p *RefundPlan) (*payment. }) return nil, err } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.Refund(ctx, payment.RefundRequest{ TradeNo: p.Order.PaymentTradeNo, OrderID: p.Order.OutTradeNo, Amount: formatGatewayRefundAmount(p.GatewayAmount, p.Order), Reason: p.Reason, }) + finishProviderCall() if err != nil { if resp != nil && strings.TrimSpace(resp.Status) == payment.ProviderStatusPending { return resp, nil @@ -417,12 +420,14 @@ func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64) } pendingDetail := s.latestRefundPendingDetail(ctx, oid) + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := queryProvider.QueryRefund(ctx, payment.RefundQueryRequest{ TradeNo: o.PaymentTradeNo, OrderID: o.OutTradeNo, RefundID: pendingDetail.RefundID, Amount: formatGatewayRefundAmount(o.RefundAmount, o), }) + finishProviderCall() if err != nil { return nil, fmt.Errorf("query refund: %w", err) } diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 100b240785..507e02d877 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -2006,9 +2006,19 @@ func parseOpenAIImageTryAgainCooldown(body []byte) time.Duration { const upstreamModelNotFoundCooldown = 30 * time.Minute const upstreamModelNotFoundReason = "upstream_404_model_not_found" +const upstreamCodexPlanGatedModelCooldown = 30 * time.Minute +const upstreamCodexPlanGatedModelReason = "upstream_400_codex_plan_gated_model" const tempUnschedBodyMaxBytes = 64 << 10 const tempUnschedMessageMaxBytes = 2048 +// HandleUpstreamModelNotFound marks the requested model as temporarily +// unavailable on the account when the upstream deterministically reports it +// cannot serve that model: a 404 model-not-found, or the Codex 400 rejecting a +// plan-gated model on a ChatGPT OAuth account. Returning true tells the caller +// to fail the current attempt over to another account; the scheduler skips the +// (account, model) pair via IsSchedulableForModelWithContext until the +// cooldown expires, instead of re-selecting an account that can never serve +// the model. func (s *RateLimitService) HandleUpstreamModelNotFound(ctx context.Context, account *Account, requestedModel string, statusCode int, responseBody []byte) bool { if s == nil || account == nil || s.accountRepo == nil { return false @@ -2016,19 +2026,26 @@ func (s *RateLimitService) HandleUpstreamModelNotFound(ctx context.Context, acco if !account.ShouldHandleErrorCode(statusCode) { return false } - if !isUpstreamModelNotFoundError(statusCode, responseBody) { + var cooldown time.Duration + var reason string + switch { + case isUpstreamModelNotFoundError(statusCode, responseBody): + cooldown, reason = upstreamModelNotFoundCooldown, upstreamModelNotFoundReason + case isOpenAIOAuthAccount(account) && isOpenAICodexPlanGatedModelError(statusCode, responseBody): + cooldown, reason = upstreamCodexPlanGatedModelCooldown, upstreamCodexPlanGatedModelReason + default: return false } modelKey := modelRateLimitKeyForUpstreamModelNotFound(ctx, account, requestedModel) if modelKey == "" { return false } - resetAt := time.Now().Add(upstreamModelNotFoundCooldown) - if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, resetAt, upstreamModelNotFoundReason); err != nil { - slog.Warn("upstream_model_not_found_set_model_rate_limit_failed", "account_id", account.ID, "model", modelKey, "error", err) + resetAt := time.Now().Add(cooldown) + if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, resetAt, reason); err != nil { + slog.Warn("upstream_model_not_found_set_model_rate_limit_failed", "account_id", account.ID, "model", modelKey, "reason", reason, "error", err) return true } - slog.Info("upstream_model_not_found_model_rate_limited", "account_id", account.ID, "model", modelKey, "reset_at", resetAt) + slog.Info("upstream_model_not_found_model_rate_limited", "account_id", account.ID, "model", modelKey, "reason", reason, "reset_at", resetAt) return true } diff --git a/backend/internal/service/ratelimit_service_model_not_found_test.go b/backend/internal/service/ratelimit_service_model_not_found_test.go index dfd18c5f69..51bd8a607e 100644 --- a/backend/internal/service/ratelimit_service_model_not_found_test.go +++ b/backend/internal/service/ratelimit_service_model_not_found_test.go @@ -125,3 +125,77 @@ func openAIModelNotFoundTempAccount() *Account { }, } } + +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedModelUsesModelRateLimit(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-5.6-sol", + ) + + require.True(t, handled) + require.Zero(t, repo.tempCalls) + require.Len(t, repo.modelRateLimitCalls, 1) + call := repo.modelRateLimitCalls[0] + require.Equal(t, account.ID, call.accountID) + require.Equal(t, "gpt-5.6-sol", call.scope) + require.Equal(t, upstreamCodexPlanGatedModelReason, call.reason) + require.WithinDuration(t, time.Now().Add(upstreamCodexPlanGatedModelCooldown), call.resetAt, 5*time.Second) +} + +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedModelRespectsModelMapping(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + account.Credentials["model_mapping"] = map[string]any{"gpt-5.6-sol": "gpt-5.6-sol-upstream"} + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-5.6-sol-upstream' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-5.6-sol", + ) + + require.True(t, handled) + require.Len(t, repo.modelRateLimitCalls, 1) + require.Equal(t, "gpt-5.6-sol-upstream", repo.modelRateLimitCalls[0].scope) +} + +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedModelIgnoresAPIKeyAccount(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + account.Type = AccountTypeAPIKey + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-5.6-sol", + ) + + require.False(t, handled) + require.Empty(t, repo.modelRateLimitCalls) +} + +func openAICodexPlanGatedOAuthAccount() *Account { + return &Account{ + ID: 202, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{}, + } +} diff --git a/backend/internal/service/scheduler_outbox.go b/backend/internal/service/scheduler_outbox.go index 2b7665ad78..a44f2a3a30 100644 --- a/backend/internal/service/scheduler_outbox.go +++ b/backend/internal/service/scheduler_outbox.go @@ -17,6 +17,8 @@ type SchedulerOutboxEvent struct { // SchedulerOutboxRepository 提供调度 outbox 的读取接口。 type SchedulerOutboxRepository interface { ListAfterAndReleaseDedup(ctx context.Context, afterID int64, limit int) ([]SchedulerOutboxEvent, error) + // FirstCreatedAtAfter 返回指定水位之后第一条待消费事件的创建时间,不领取事件或修改去重键。 + FirstCreatedAtAfter(ctx context.Context, afterID int64) (time.Time, bool, error) MaxID(ctx context.Context) (int64, error) DeleteConsumedUpTo(ctx context.Context, watermark int64, limit int) (int64, error) TryAcquireCleanupLock(ctx context.Context) (SchedulerOutboxCleanupLease, bool, error) diff --git a/backend/internal/service/scheduler_snapshot_full_rebuild_test.go b/backend/internal/service/scheduler_snapshot_full_rebuild_test.go new file mode 100644 index 0000000000..09ec2790ce --- /dev/null +++ b/backend/internal/service/scheduler_snapshot_full_rebuild_test.go @@ -0,0 +1,145 @@ +package service + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type schedulerFullRebuildTestCache struct { + SchedulerCache + + mu sync.Mutex + listErr error + listCalls int + lockCalls int +} + +func (c *schedulerFullRebuildTestCache) ListBuckets(context.Context) ([]SchedulerBucket, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.listCalls++ + return nil, c.listErr +} + +func (c *schedulerFullRebuildTestCache) TryLockBucket(context.Context, SchedulerBucket, time.Duration) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.lockCalls++ + return false, nil +} + +func TestSchedulerSnapshotServiceFullRebuildCoalescesConcurrentRequestsIntoTrailingRun(t *testing.T) { + svc := &SchedulerSnapshotService{} + wantTrailingErr := errors.New("trailing rebuild failed") + firstStarted := make(chan struct{}) + releaseFirst := make(chan struct{}) + var releaseOnce sync.Once + release := func() { + releaseOnce.Do(func() { close(releaseFirst) }) + } + defer release() + + var calls atomic.Int32 + var active atomic.Int32 + var maxActive atomic.Int32 + run := func() error { + call := calls.Add(1) + currentActive := active.Add(1) + defer active.Add(-1) + for { + previousMax := maxActive.Load() + if currentActive <= previousMax || maxActive.CompareAndSwap(previousMax, currentActive) { + break + } + } + if call == 1 { + close(firstStarted) + <-releaseFirst + return nil + } + return wantTrailingErr + } + + firstResult := make(chan error, 1) + go func() { + firstResult <- svc.coalesceFullRebuild(run) + }() + + select { + case <-firstStarted: + case <-time.After(time.Second): + t.Fatal("first rebuild did not start") + } + + const followers = 20 + followerResults := make(chan error, followers) + for range followers { + go func() { + followerResults <- svc.coalesceFullRebuild(run) + }() + } + + require.Eventually(t, func() bool { + requested, _ := schedulerFullRebuildState(svc) + return requested == followers+1 + }, time.Second, time.Millisecond) + release() + + require.NoError(t, <-firstResult) + for range followers { + require.ErrorIs(t, <-followerResults, wantTrailingErr) + } + require.EqualValues(t, 2, calls.Load()) + require.EqualValues(t, 1, maxActive.Load()) + requested, completed := schedulerFullRebuildState(svc) + require.EqualValues(t, followers+1, requested) + require.Equal(t, requested, completed) +} + +func TestSchedulerSnapshotServiceFullRebuildRunsAgainForSequentialRequest(t *testing.T) { + svc := &SchedulerSnapshotService{} + wantSecondErr := errors.New("second rebuild failed") + var calls atomic.Int32 + run := func() error { + if calls.Add(1) == 2 { + return wantSecondErr + } + return nil + } + + require.NoError(t, svc.coalesceFullRebuild(run)) + require.ErrorIs(t, svc.coalesceFullRebuild(run), wantSecondErr) + require.EqualValues(t, 2, calls.Load()) + requested, completed := schedulerFullRebuildState(svc) + require.EqualValues(t, 2, requested) + require.Equal(t, requested, completed) +} + +func TestSchedulerSnapshotServiceInitialFullRebuildFallsBackWhenListBucketsFails(t *testing.T) { + cache := &schedulerFullRebuildTestCache{listErr: errors.New("list buckets failed")} + svc := NewSchedulerSnapshotService(cache, nil, nil, nil, nil) + + svc.runInitialRebuild() + + cache.mu.Lock() + listCalls := cache.listCalls + lockCalls := cache.lockCalls + cache.mu.Unlock() + require.Equal(t, 1, listCalls) + require.Positive(t, lockCalls, "startup should rebuild default buckets after ListBuckets fails") + requested, completed := schedulerFullRebuildState(svc) + require.EqualValues(t, 1, requested) + require.Equal(t, requested, completed) +} + +func schedulerFullRebuildState(svc *SchedulerSnapshotService) (requested uint64, completed uint64) { + svc.fullRebuildStateMu.Lock() + defer svc.fullRebuildStateMu.Unlock() + return svc.fullRebuildRequested, svc.fullRebuildCompleted +} diff --git a/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go b/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go index 535f8d54e2..91e2d36f76 100644 --- a/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go +++ b/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go @@ -6,12 +6,15 @@ import ( "reflect" "testing" "time" + + "github.com/Wei-Shaw/sub2api/internal/config" ) type outboxCleanupCache struct { - watermark int64 - setWatermarks []int64 - updateErr error + watermark int64 + setWatermarks []int64 + updateErr error + listBucketCalls int } func (c *outboxCleanupCache) GetSnapshot(ctx context.Context, bucket SchedulerBucket) ([]*Account, bool, error) { @@ -47,6 +50,7 @@ func (c *outboxCleanupCache) UnlockBucket(ctx context.Context, bucket SchedulerB } func (c *outboxCleanupCache) ListBuckets(ctx context.Context) ([]SchedulerBucket, error) { + c.listBucketCalls++ return nil, nil } @@ -66,12 +70,13 @@ type outboxCleanupDeleteCall struct { } type outboxCleanupRepo struct { - events []SchedulerOutboxEvent - rows []int64 - lockAcquired bool - lockAttempts int - releaseCount int - deleteCalls []outboxCleanupDeleteCall + events []SchedulerOutboxEvent + rows []int64 + lockAcquired bool + lockAttempts int + releaseCount int + deleteCalls []outboxCleanupDeleteCall + firstCreatedAfterID []int64 } func (r *outboxCleanupRepo) ListAfterAndReleaseDedup(ctx context.Context, afterID int64, limit int) ([]SchedulerOutboxEvent, error) { @@ -88,6 +93,16 @@ func (r *outboxCleanupRepo) ListAfterAndReleaseDedup(ctx context.Context, afterI return events, nil } +func (r *outboxCleanupRepo) FirstCreatedAtAfter(ctx context.Context, afterID int64) (time.Time, bool, error) { + r.firstCreatedAfterID = append(r.firstCreatedAfterID, afterID) + for _, event := range r.events { + if event.ID > afterID { + return event.CreatedAt, true, nil + } + } + return time.Time{}, false, nil +} + func (r *outboxCleanupRepo) MaxID(ctx context.Context) (int64, error) { var maxID int64 for _, id := range r.rows { @@ -240,6 +255,44 @@ func TestSchedulerSnapshotServicePollOutboxDoesNotCleanupOnHandleFailure(t *test } } +func TestSchedulerSnapshotServicePollOutboxDoesNotUseConsumedEventForLag(t *testing.T) { + cache := &outboxCleanupCache{} + repo := &outboxCleanupRepo{ + events: []SchedulerOutboxEvent{ + { + ID: 7, + EventType: SchedulerOutboxEventAccountLastUsed, + CreatedAt: time.Now().Add(-time.Hour), + }, + }, + } + cfg := &config.Config{ + Gateway: config.GatewayConfig{ + Scheduling: config.GatewaySchedulingConfig{ + OutboxLagWarnSeconds: 1, + OutboxLagRebuildSeconds: 1, + OutboxLagRebuildFailures: 1, + }, + }, + } + svc := NewSchedulerSnapshotService(cache, repo, nil, nil, cfg) + + svc.pollOutbox() + + if cache.watermark != 7 { + t.Fatalf("expected watermark 7, got %d", cache.watermark) + } + if !reflect.DeepEqual(repo.firstCreatedAfterID, []int64{7}) { + t.Fatalf("expected lag check after consumed watermark, got %#v", repo.firstCreatedAfterID) + } + if cache.listBucketCalls != 0 { + t.Fatalf("expected consumed event not to trigger full rebuild, got %d attempts", cache.listBucketCalls) + } + if svc.lagFailures != 0 { + t.Fatalf("expected lag failures to remain reset, got %d", svc.lagFailures) + } +} + func TestSchedulerSnapshotServiceCleanupSkipsNonPositiveWatermark(t *testing.T) { repo := &outboxCleanupRepo{ rows: []int64{1, 2, 3}, diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index dc514bc851..77bef26edf 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -43,6 +43,12 @@ type SchedulerSnapshotService struct { fallbackLimit *fallbackLimiter lagMu sync.Mutex lagFailures int + + fullRebuildRunMu sync.Mutex + fullRebuildStateMu sync.Mutex + fullRebuildRequested uint64 + fullRebuildCompleted uint64 + fullRebuildLastErr error } func NewSchedulerSnapshotService( @@ -183,22 +189,26 @@ func (s *SchedulerSnapshotService) runInitialRebuild() { if s.cache == nil { return } - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) - defer cancel() - buckets, err := s.cache.ListBuckets(ctx) - if err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] list buckets failed: %v", err) - } - if len(buckets) == 0 { - buckets, err = s.defaultBuckets(ctx) + _ = s.coalesceFullRebuild(func() error { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + buckets, err := s.cache.ListBuckets(ctx) if err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] default buckets failed: %v", err) - return + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] list buckets failed: %v", err) } - } - if err := s.rebuildBuckets(ctx, buckets, "startup"); err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] rebuild startup failed: %v", err) - } + if len(buckets) == 0 { + buckets, err = s.defaultBuckets(ctx) + if err != nil { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] default buckets failed: %v", err) + return err + } + } + if err := s.rebuildBuckets(ctx, buckets, "startup"); err != nil { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] rebuild startup failed: %v", err) + return err + } + return nil + }) } func (s *SchedulerSnapshotService) runOutboxWorker(interval time.Duration) { @@ -254,7 +264,6 @@ func (s *SchedulerSnapshotService) pollOutbox() { return } - watermarkForCheck := watermark seen := make(map[batchSeenKey]struct{}) for _, event := range events { eventCtx, cancel := context.WithTimeout(context.Background(), outboxEventTimeout) @@ -281,12 +290,15 @@ func (s *SchedulerSnapshotService) pollOutbox() { } if wmErr != nil { logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox watermark write failed: %v", wmErr) - } else { - watermarkForCheck = lastID - s.cleanupConsumedOutbox(lastID) + return } + s.cleanupConsumedOutbox(lastID) - s.checkOutboxLag(ctx, events[0], watermarkForCheck) + // 只有 watermark 成功推进后,当前批次才算已消费。延迟必须按下一条待消费事件计算, + // 否则本批次处理越慢,越容易误触发一次更慢的全量重建,形成正反馈。 + lagCtx, lagCancel := context.WithTimeout(context.Background(), 5*time.Second) + s.checkOutboxLag(lagCtx, lastID) + lagCancel() } func (s *SchedulerSnapshotService) cleanupConsumedOutbox(watermark int64) { @@ -602,30 +614,72 @@ func (s *SchedulerSnapshotService) triggerFullRebuild(reason string) error { if s.cache == nil { return ErrSchedulerCacheNotReady } - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) - defer cancel() + return s.coalesceFullRebuild(func() error { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() - buckets, err := s.cache.ListBuckets(ctx) - if err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] list buckets failed: %v", err) - return err - } - if len(buckets) == 0 { - buckets, err = s.defaultBuckets(ctx) + buckets, err := s.cache.ListBuckets(ctx) if err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] default buckets failed: %v", err) + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] list buckets failed: %v", err) return err } - } - return s.rebuildBuckets(ctx, buckets, reason) + if len(buckets) == 0 { + buckets, err = s.defaultBuckets(ctx) + if err != nil { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] default buckets failed: %v", err) + return err + } + } + return s.rebuildBuckets(ctx, buckets, reason) + }) } -func (s *SchedulerSnapshotService) checkOutboxLag(ctx context.Context, oldest SchedulerOutboxEvent, watermark int64) { - if oldest.CreatedAt.IsZero() || s.cfg == nil { +func (s *SchedulerSnapshotService) coalesceFullRebuild(run func() error) error { + s.fullRebuildStateMu.Lock() + s.fullRebuildRequested++ + requestID := s.fullRebuildRequested + s.fullRebuildStateMu.Unlock() + + s.fullRebuildRunMu.Lock() + defer s.fullRebuildRunMu.Unlock() + + s.fullRebuildStateMu.Lock() + if s.fullRebuildCompleted >= requestID { + err := s.fullRebuildLastErr + s.fullRebuildStateMu.Unlock() + return err + } + // 当前轮重建可能早于新 outbox 事件对应事务的提交,不能让后到请求直接复用当前轮。 + // 每轮开始前记录可覆盖的请求代次,执行期间登记的请求统一合并到下一轮。 + coveredThrough := s.fullRebuildRequested + s.fullRebuildStateMu.Unlock() + + err := run() + + s.fullRebuildStateMu.Lock() + s.fullRebuildCompleted = coveredThrough + s.fullRebuildLastErr = err + s.fullRebuildStateMu.Unlock() + return err +} + +func (s *SchedulerSnapshotService) checkOutboxLag(ctx context.Context, watermark int64) { + if s.cfg == nil || s.outboxRepo == nil { + return + } + oldestCreatedAt, ok, err := s.outboxRepo.FirstCreatedAtAfter(ctx, watermark) + if err != nil { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox pending event read failed: %v", err) + return + } + if !ok || oldestCreatedAt.IsZero() { + s.lagMu.Lock() + s.lagFailures = 0 + s.lagMu.Unlock() return } - lag := time.Since(oldest.CreatedAt) + lag := time.Since(oldestCreatedAt) if lagSeconds := int(lag.Seconds()); lagSeconds >= s.cfg.Gateway.Scheduling.OutboxLagWarnSeconds && s.cfg.Gateway.Scheduling.OutboxLagWarnSeconds > 0 { logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox lag warning: %ds", lagSeconds) } @@ -652,7 +706,7 @@ func (s *SchedulerSnapshotService) checkOutboxLag(ctx context.Context, oldest Sc } threshold := s.cfg.Gateway.Scheduling.OutboxBacklogRebuildRows - if threshold <= 0 || s.outboxRepo == nil { + if threshold <= 0 { return } maxID, err := s.outboxRepo.MaxID(ctx) diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index 4f6a305b25..9b4fd6f532 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -131,6 +131,8 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc switch { case account.Platform == PlatformAntigravity: return s.buildAntigravityAPIKeyModelsRequest(ctx, account) + case account.IsGrok(): + return s.buildGrokUpstreamModelsRequest(ctx, account) case account.IsOpenAI(): return s.buildOpenAIUpstreamModelsRequest(ctx, account) case account.IsGemini(): @@ -144,6 +146,36 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc } } +func (s *AccountTestService) buildGrokUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) { + if account.Type != AccountTypeAPIKey { + return nil, newUpstreamModelSyncUnsupportedError( + fmt.Sprintf("Unsupported Grok account type for upstream model sync: %s", account.Type), nil, + ) + } + apiKey := strings.TrimSpace(account.GetCredential("api_key")) + if apiKey == "" { + return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil) + } + + baseURL := strings.TrimSpace(account.GetCredential("base_url")) + if baseURL == "" { + baseURL = "https://api.x.ai" + } + normalizedBaseURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, buildOpenAIModelsURL(normalizedBaseURL), nil) + if err != nil { + return nil, newUpstreamModelSyncConfigError("Invalid Grok model list URL", err) + } + req.Header.Set("Accept", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + account.ApplyHeaderOverrides(req.Header) + return req, nil +} + func (s *AccountTestService) buildAnthropicUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) { if account.IsBedrock() || account.Type == AccountTypeServiceAccount { return nil, newUpstreamModelSyncUnsupportedError( diff --git a/backend/internal/service/upstream_models_test.go b/backend/internal/service/upstream_models_test.go index 3904194ffa..5b5c5e9835 100644 --- a/backend/internal/service/upstream_models_test.go +++ b/backend/internal/service/upstream_models_test.go @@ -177,6 +177,18 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) { require.Equal(t, "https://openai.example.com/v1/models", openAIReq.URL.String()) require.Equal(t, "Bearer openai-key", openAIReq.Header.Get("Authorization")) + grokReq, err := svc.buildUpstreamModelsRequest(ctx, &Account{ + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "xai-key", + "base_url": "https://xai.example.com/v1", + }, + }) + require.NoError(t, err) + require.Equal(t, "https://xai.example.com/v1/models", grokReq.URL.String()) + require.Equal(t, "Bearer xai-key", grokReq.Header.Get("Authorization")) + geminiReq, err := svc.buildGeminiUpstreamModelsRequest(ctx, &Account{ Platform: PlatformGemini, Type: AccountTypeAPIKey, @@ -202,6 +214,22 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) { require.Equal(t, "antigravity-key", antigravityReq.Header.Get("x-api-key")) } +func TestBuildUpstreamModelsRequestRejectsGrokOAuth(t *testing.T) { + t.Parallel() + + svc := &AccountTestService{cfg: upstreamModelSyncTestConfig()} + _, err := svc.buildUpstreamModelsRequest(context.Background(), &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + }) + require.Error(t, err) + + var syncErr *UpstreamModelSyncError + require.True(t, errors.As(err, &syncErr)) + require.Equal(t, UpstreamModelSyncErrorUnsupported, syncErr.Kind) + require.Contains(t, syncErr.SafeMessage(), "Unsupported Grok account type") +} + func TestBuildAntigravityAPIKeyModelsRequestRejectsOfficialCloudCodeBase(t *testing.T) { t.Parallel() @@ -265,6 +293,34 @@ func TestFetchUpstreamSupportedModelsParsesOpenAIResponse(t *testing.T) { require.Equal(t, "Bearer openai-key", upstream.lastReq.Header.Get("Authorization")) } +func TestFetchUpstreamSupportedModelsParsesGrokAPIKeyResponse(t *testing.T) { + t.Parallel() + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"grok-4.5"},{"id":"grok-4.5"},{"id":"grok-imagine"}]}`)), + }} + svc := &AccountTestService{ + httpUpstream: upstream, + cfg: upstreamModelSyncTestConfig(), + } + + models, err := svc.FetchUpstreamSupportedModels(context.Background(), &Account{ + ID: 9, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "xai-key", + "base_url": "https://xai.example.com/v1", + }, + }) + require.NoError(t, err) + require.Equal(t, []string{"grok-4.5", "grok-imagine"}, models) + require.Equal(t, "https://xai.example.com/v1/models", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer xai-key", upstream.lastReq.Header.Get("Authorization")) +} + func TestFetchUpstreamSupportedModelsDoesNotExposeUpstreamBody(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index 62e48fc8f9..0adcc04a94 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -142,13 +142,14 @@ type UsageLog struct { ImageOutputTokens int ImageOutputCost float64 - InputCost float64 - OutputCost float64 - CacheCreationCost float64 - CacheReadCost float64 - TotalCost float64 - ActualCost float64 - RateMultiplier float64 + InputCost float64 + OutputCost float64 + CacheCreationCost float64 + CacheReadCost float64 + TotalCost float64 + ActualCost float64 + RateMultiplier float64 + LongContextBillingApplied bool // AccountRateMultiplier 账号计费倍率快照(nil 表示历史数据,按 1.0 处理) AccountRateMultiplier *float64 // AccountStatsCost 账号统计定价预计算费用(nil = 使用默认公式 total_cost × account_rate_multiplier) diff --git a/backend/internal/service/vertex_service_account.go b/backend/internal/service/vertex_service_account.go index 256695ded5..7ccbeee43c 100644 --- a/backend/internal/service/vertex_service_account.go +++ b/backend/internal/service/vertex_service_account.go @@ -18,6 +18,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/golang-jwt/jwt/v5" ) @@ -195,7 +196,7 @@ func vertexServiceAccountProxyURL(account *Account) string { func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) { proxyURL = strings.TrimSpace(proxyURL) if proxyURL == "" { - return &http.Client{Timeout: 15 * time.Second}, nil + return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second}), nil } _, parsedProxy, err := proxyurl.Parse(proxyURL) @@ -211,7 +212,7 @@ func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) { if err := proxyutil.ConfigureTransportProxy(transport, parsedProxy); err != nil { return nil, err } - return &http.Client{Timeout: 15 * time.Second, Transport: transport}, nil + return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second, Transport: transport}), nil } func exchangeVertexServiceAccountToken(ctx context.Context, key *vertexServiceAccountKey, proxyURL string) (string, time.Duration, error) { diff --git a/backend/internal/service/vertex_service_account_test.go b/backend/internal/service/vertex_service_account_test.go index d77a1988e9..68c756eaa2 100644 --- a/backend/internal/service/vertex_service_account_test.go +++ b/backend/internal/service/vertex_service_account_test.go @@ -13,6 +13,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -101,6 +102,24 @@ func TestVertexServiceAccountProxyURL(t *testing.T) { require.Empty(t, vertexServiceAccountProxyURL(&Account{ProxyID: &proxyID})) } +func TestVertexServiceAccountHTTPClientRecordsDependency(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + client, err := newVertexServiceAccountHTTPClient("") + require.NoError(t, err) + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + request, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil) + require.NoError(t, err) + response, err := client.Do(request) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + require.Contains(t, collector.HeaderValue(time.Now(), "bypass"), "dep_http;dur=") +} + func TestExchangeVertexServiceAccountTokenUsesProxy(t *testing.T) { privateKey, err := rsa.GenerateKey(rand.Reader, 2048) require.NoError(t, err) diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 8c20af661b..8dad4b51a2 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -145,6 +145,7 @@ func ProvideAccountUsageService( geminiQuotaService *GeminiQuotaService, antigravityQuotaFetcher *AntigravityQuotaFetcher, grokQuotaFetcher *GrokQuotaFetcher, + grokQuotaService *GrokQuotaService, openAIQuotaService *OpenAIQuotaService, cache *UsageCache, identityCache IdentityCache, @@ -158,6 +159,7 @@ func ProvideAccountUsageService( geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, + grokQuotaService, openAIQuotaService, cache, identityCache, @@ -197,8 +199,9 @@ func ProvideGrokQuotaService( proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, + usageLogRepo UsageLogRepository, ) *GrokQuotaService { - return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream) + return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream, usageLogRepo) } // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection diff --git a/backend/internal/web/embed_on.go b/backend/internal/web/embed_on.go index 41738e7a5d..716fb77e75 100644 --- a/backend/internal/web/embed_on.go +++ b/backend/internal/web/embed_on.go @@ -109,7 +109,8 @@ func (s *FrontendServer) Middleware() gin.HandlerFunc { return } - // Serve static files normally + // Serve static files normally (hashed assets get long-lived cache headers) + applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath) s.fileServer.ServeHTTP(c.Writer, c.Request) c.Abort() } @@ -135,6 +136,7 @@ func (s *FrontendServer) tryServeOverride(c *gin.Context, cleanPath string) bool if err != nil || info.IsDir() { return false } + applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath) c.File(filePath) c.Abort() return true @@ -273,6 +275,7 @@ func ServeEmbeddedFrontend() gin.HandlerFunc { if tryServeOverrideFile(c, overrideDir, cleanPath) { return } + applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath) fileServer.ServeHTTP(c.Writer, c.Request) c.Abort() return @@ -292,6 +295,7 @@ func tryServeOverrideFile(c *gin.Context, overrideDir, cleanPath string) bool { if err != nil || info.IsDir() { return false } + applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath) c.File(filePath) c.Abort() return true @@ -308,7 +312,9 @@ func shouldBypassEmbeddedFrontend(path string) bool { trimmed == "/health" || trimmed == "/responses" || strings.HasPrefix(trimmed, "/responses/") || - strings.HasPrefix(trimmed, "/images/") + trimmed == "/alpha/search" || + strings.HasPrefix(trimmed, "/images/") || + strings.HasPrefix(trimmed, "/videos/") } func serveIndexHTML(c *gin.Context, fsys fs.FS) { diff --git a/backend/internal/web/embed_test.go b/backend/internal/web/embed_test.go index 27e15ef166..b27bbfc9dc 100644 --- a/backend/internal/web/embed_test.go +++ b/backend/internal/web/embed_test.go @@ -507,6 +507,32 @@ func TestFrontendServer_Middleware(t *testing.T) { assert.JSONEq(t, `{"ok":true}`, w.Body.String()) }) + t.Run("skips_alpha_search_post_route", func(t *testing.T) { + provider := &mockSettingsProvider{ + settings: map[string]string{"test": "value"}, + } + + server, err := NewFrontendServer(provider) + require.NoError(t, err) + + router := gin.New() + router.Use(server.Middleware()) + nextCalled := false + router.POST("/alpha/search", func(c *gin.Context) { + nextCalled = true + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + assert.True(t, nextCalled, "next handler should be called for alpha search API route") + assert.Equal(t, http.StatusOK, w.Code) + assert.JSONEq(t, `{"ok":true}`, w.Body.String()) + }) + t.Run("serves_index_for_spa_routes", func(t *testing.T) { provider := &mockSettingsProvider{ settings: map[string]string{"test": "value"}, @@ -562,6 +588,17 @@ func TestFrontendServer_Middleware(t *testing.T) { }) } +func TestEmbeddedFrontendBypassesBareVideoAPIRoutes(t *testing.T) { + for _, path := range []string{ + "/videos/generations", + "/videos/edits", + "/videos/extensions", + "/videos/request-123", + } { + require.True(t, shouldBypassEmbeddedFrontend(path), "path=%s", path) + } +} + func TestNewFrontendServer(t *testing.T) { t.Run("creates_server_successfully", func(t *testing.T) { provider := &mockSettingsProvider{ diff --git a/backend/internal/web/static_cache.go b/backend/internal/web/static_cache.go new file mode 100644 index 0000000000..09abc8642a --- /dev/null +++ b/backend/internal/web/static_cache.go @@ -0,0 +1,31 @@ +//go:build embed || unit + +package web + +import ( + "net/http" + "strings" +) + +// staticAssetsCacheControl matches deploy/Caddyfile for hashed frontend assets. +// Vite emits content-hashed filenames under assets/, so long-lived immutable +// caching is safe without relying on a reverse proxy. +const staticAssetsCacheControl = "public, max-age=31536000, immutable" + +// isLongCacheStaticPath reports whether a cleaned URL path (no leading slash) +// should receive long-lived Cache-Control headers. Aligned with deploy/Caddyfile. +func isLongCacheStaticPath(cleanPath string) bool { + cleanPath = strings.TrimPrefix(cleanPath, "/") + return strings.HasPrefix(cleanPath, "assets/") || + cleanPath == "logo.png" || + cleanPath == "favicon.ico" +} + +// applyStaticAssetCacheHeaders sets Cache-Control for long-cacheable static paths. +// index.html / SPA routes must keep no-cache and are not handled here. +func applyStaticAssetCacheHeaders(header http.Header, cleanPath string) { + if header == nil || !isLongCacheStaticPath(cleanPath) { + return + } + header.Set("Cache-Control", staticAssetsCacheControl) +} diff --git a/backend/internal/web/static_cache_test.go b/backend/internal/web/static_cache_test.go new file mode 100644 index 0000000000..130347c41e --- /dev/null +++ b/backend/internal/web/static_cache_test.go @@ -0,0 +1,71 @@ +//go:build unit + +package web + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestIsLongCacheStaticPath(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + path string + want bool + }{ + {name: "hashed_js", path: "assets/index-abc123.js", want: true}, + {name: "hashed_css", path: "assets/app-def456.css", want: true}, + {name: "nested_asset", path: "assets/vendor/chunk.js", want: true}, + {name: "leading_slash_asset", path: "/assets/index.js", want: true}, + {name: "logo", path: "logo.png", want: true}, + {name: "favicon", path: "favicon.ico", want: true}, + {name: "index_html", path: "index.html", want: false}, + {name: "spa_route", path: "dashboard", want: false}, + {name: "assets_prefix_only", path: "assets", want: false}, + {name: "similar_name", path: "assets-backup/x.js", want: false}, + {name: "empty", path: "", want: false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tc.want, isLongCacheStaticPath(tc.path)) + }) + } +} + +func TestApplyStaticAssetCacheHeaders(t *testing.T) { + t.Parallel() + + t.Run("sets_immutable_cache_for_assets", func(t *testing.T) { + t.Parallel() + header := make(http.Header) + applyStaticAssetCacheHeaders(header, "assets/index-abc.js") + assert.Equal(t, staticAssetsCacheControl, header.Get("Cache-Control")) + }) + + t.Run("sets_immutable_cache_for_logo", func(t *testing.T) { + t.Parallel() + header := make(http.Header) + applyStaticAssetCacheHeaders(header, "logo.png") + assert.Equal(t, staticAssetsCacheControl, header.Get("Cache-Control")) + }) + + t.Run("skips_index_html", func(t *testing.T) { + t.Parallel() + header := make(http.Header) + applyStaticAssetCacheHeaders(header, "index.html") + assert.Empty(t, header.Get("Cache-Control")) + }) + + t.Run("nil_header_is_noop", func(t *testing.T) { + t.Parallel() + assert.NotPanics(t, func() { + applyStaticAssetCacheHeaders(nil, "assets/x.js") + }) + }) +} diff --git a/backend/migrations/174_add_usage_log_long_context_billing.sql b/backend/migrations/174_add_usage_log_long_context_billing.sql new file mode 100644 index 0000000000..090403c310 --- /dev/null +++ b/backend/migrations/174_add_usage_log_long_context_billing.sql @@ -0,0 +1,4 @@ +-- Snapshot whether long-context pricing changed token prices for a request so +-- usage history can explain the applied charge without inferring from totals. +ALTER TABLE usage_logs + ADD COLUMN IF NOT EXISTS long_context_billing_applied BOOLEAN NOT NULL DEFAULT FALSE; diff --git a/backend/migrations/174_add_usage_logs_api_key_latest_ip_index_notx.sql b/backend/migrations/174_add_usage_logs_api_key_latest_ip_index_notx.sql new file mode 100644 index 0000000000..261698f8cc --- /dev/null +++ b/backend/migrations/174_add_usage_logs_api_key_latest_ip_index_notx.sql @@ -0,0 +1,5 @@ +-- Support the per-key latest non-empty source IP lookup without scanning full key history. +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip + ON usage_logs (api_key_id, created_at DESC, id DESC) + INCLUDE (ip_address) + WHERE ip_address IS NOT NULL AND ip_address <> ''; diff --git a/backend/migrations/174_group_web_search_price_per_call.sql b/backend/migrations/174_group_web_search_price_per_call.sql new file mode 100644 index 0000000000..da9c90fed8 --- /dev/null +++ b/backend/migrations/174_group_web_search_price_per_call.sql @@ -0,0 +1,3 @@ +-- Codex alpha/search 网页搜索按次计费:分组级单次价格覆盖。 +-- NULL 表示使用内置默认价 0.01 USD/次(OpenAI 官方 web search 定价 $10/1000 次)。 +ALTER TABLE groups ADD COLUMN IF NOT EXISTS web_search_price_per_call DECIMAL(20,8); diff --git a/backend/migrations/175_add_ops_system_logs_host.sql b/backend/migrations/175_add_ops_system_logs_host.sql new file mode 100644 index 0000000000..e5f9f7299c --- /dev/null +++ b/backend/migrations/175_add_ops_system_logs_host.sql @@ -0,0 +1,3 @@ +-- Track the application host that emitted each indexed system log. +ALTER TABLE ops_system_logs + ADD COLUMN IF NOT EXISTS host VARCHAR(255); diff --git a/backend/migrations/175_default_openai_long_context_billing.sql b/backend/migrations/175_default_openai_long_context_billing.sql new file mode 100644 index 0000000000..cccbea4108 --- /dev/null +++ b/backend/migrations/175_default_openai_long_context_billing.sql @@ -0,0 +1,162 @@ +-- Keep mixed-version writers consistent before backfilling rows that already exist. +CREATE OR REPLACE FUNCTION public.enforce_openai_long_context_billing_extra() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +DECLARE + parent_effective_value JSONB; +BEGIN + IF NEW.platform IS DISTINCT FROM 'openai' THEN + RETURN NEW; + END IF; + + NEW.extra := COALESCE(NEW.extra, '{}'::jsonb); + IF NEW.parent_account_id IS NOT NULL AND NEW.quota_dimension = 'spark' THEN + SELECT CASE + WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb + WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb + WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean' + THEN parent.extra->'openai_long_context_billing_enabled' + ELSE 'false'::jsonb + END + INTO parent_effective_value + FROM accounts AS parent + WHERE parent.id = NEW.parent_account_id; + + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + COALESCE(parent_effective_value, 'false'::jsonb), + true + ); + ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled') + AND TG_OP = 'UPDATE' + AND OLD.platform = 'openai' + AND jsonb_typeof(OLD.extra->'openai_long_context_billing_enabled') = 'boolean' THEN + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + OLD.extra->'openai_long_context_billing_enabled', + true + ); + ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled') THEN + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true + ); + END IF; + + IF jsonb_typeof(NEW.extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean' THEN + RAISE EXCEPTION 'openai_long_context_billing_enabled must be a boolean' + USING ERRCODE = '22023'; + END IF; + RETURN NEW; +END; +$$; + +CREATE OR REPLACE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +BEGIN + WITH updated_shadows AS ( + UPDATE accounts AS shadow + SET extra = jsonb_set( + COALESCE(shadow.extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + NEW.extra->'openai_long_context_billing_enabled', + true + ) + WHERE shadow.parent_account_id = NEW.id + AND shadow.platform = 'openai' + AND shadow.quota_dimension = 'spark' + AND shadow.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled' + RETURNING shadow.id + ) + INSERT INTO scheduler_outbox (event_type, account_id) + SELECT 'account_changed', id + FROM updated_shadows; + RETURN NULL; +END; +$$; + +DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts; +CREATE TRIGGER accounts_enforce_openai_long_context_billing_extra +BEFORE INSERT OR UPDATE OF platform, extra, parent_account_id, quota_dimension +ON accounts +FOR EACH ROW +EXECUTE FUNCTION public.enforce_openai_long_context_billing_extra(); + +DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts; +CREATE TRIGGER accounts_propagate_openai_long_context_billing_extra +AFTER UPDATE OF platform, extra +ON accounts +FOR EACH ROW +WHEN ( + NEW.platform = 'openai' + AND NEW.parent_account_id IS NULL + AND ( + OLD.platform IS DISTINCT FROM NEW.platform + OR OLD.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled' + ) +) +EXECUTE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows(); + +UPDATE accounts +SET extra = jsonb_set( + COALESCE(extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true +) +WHERE platform = 'openai' + AND COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled' + AND jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean'; + +UPDATE accounts +SET extra = jsonb_set( + COALESCE(extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true +) +WHERE platform = 'openai' + AND parent_account_id IS NULL + AND NOT (COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled'); + +WITH shadow_values AS ( + SELECT + shadow.id, + CASE + WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb + WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb + WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean' + THEN parent.extra->'openai_long_context_billing_enabled' + ELSE 'false'::jsonb + END AS effective_value + FROM accounts AS shadow + JOIN accounts AS parent ON parent.id = shadow.parent_account_id + WHERE shadow.platform = 'openai' + AND shadow.quota_dimension = 'spark' +), +updated_shadows AS ( + UPDATE accounts AS shadow + SET extra = jsonb_set( + COALESCE(shadow.extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + shadow_values.effective_value, + true + ) + FROM shadow_values + WHERE shadow.id = shadow_values.id + AND shadow.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM shadow_values.effective_value + RETURNING shadow.id +) +INSERT INTO scheduler_outbox (event_type, account_id) +SELECT 'account_changed', id +FROM updated_shadows; diff --git a/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql new file mode 100644 index 0000000000..ec2705e49b --- /dev/null +++ b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql @@ -0,0 +1,2 @@ +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_ops_system_logs_host_created_at + ON ops_system_logs (host, created_at DESC); diff --git a/backend/migrations/176_channel_monitor_grok_provider.sql b/backend/migrations/176_channel_monitor_grok_provider.sql new file mode 100644 index 0000000000..b1bad754a4 --- /dev/null +++ b/backend/migrations/176_channel_monitor_grok_provider.sql @@ -0,0 +1,39 @@ +-- Migration: 176_channel_monitor_grok_provider +-- Allow Grok as a channel-monitor provider. Grok checks use the existing +-- OpenAI-compatible chat completions protocol with model grok-4.5 by default. + +DO $$ +DECLARE + monitor_constraint_def TEXT; + template_constraint_def TEXT; +BEGIN + SELECT pg_get_constraintdef(c.oid) + INTO monitor_constraint_def + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + WHERE t.relname = 'channel_monitors' + AND c.conname = 'channel_monitors_provider_check'; + + IF monitor_constraint_def IS NULL OR position('grok' IN monitor_constraint_def) = 0 THEN + ALTER TABLE channel_monitors + DROP CONSTRAINT IF EXISTS channel_monitors_provider_check; + ALTER TABLE channel_monitors + ADD CONSTRAINT channel_monitors_provider_check + CHECK (provider IN ('openai', 'anthropic', 'gemini', 'grok')); + END IF; + + SELECT pg_get_constraintdef(c.oid) + INTO template_constraint_def + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + WHERE t.relname = 'channel_monitor_request_templates' + AND c.conname = 'channel_monitor_request_templates_provider_check'; + + IF template_constraint_def IS NULL OR position('grok' IN template_constraint_def) = 0 THEN + ALTER TABLE channel_monitor_request_templates + DROP CONSTRAINT IF EXISTS channel_monitor_request_templates_provider_check; + ALTER TABLE channel_monitor_request_templates + ADD CONSTRAINT channel_monitor_request_templates_provider_check + CHECK (provider IN ('openai', 'anthropic', 'gemini', 'grok')); + END IF; +END $$; diff --git a/backend/migrations/channel_monitor_grok_provider_migration_test.go b/backend/migrations/channel_monitor_grok_provider_migration_test.go new file mode 100644 index 0000000000..2545173f0e --- /dev/null +++ b/backend/migrations/channel_monitor_grok_provider_migration_test.go @@ -0,0 +1,20 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestChannelMonitorGrokProviderMigration(t *testing.T) { + content, err := FS.ReadFile("176_channel_monitor_grok_provider.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "channel_monitors_provider_check") + require.Contains(t, sql, "channel_monitor_request_templates_provider_check") + require.Contains(t, sql, "CHECK (provider IN ('openai', 'anthropic', 'gemini', 'grok'))") + require.Contains(t, sql, "position('grok' IN monitor_constraint_def) = 0") + require.Contains(t, sql, "position('grok' IN template_constraint_def) = 0") +} diff --git a/backend/migrations/latest_api_key_ip_index_test.go b/backend/migrations/latest_api_key_ip_index_test.go new file mode 100644 index 0000000000..1de64a9ff5 --- /dev/null +++ b/backend/migrations/latest_api_key_ip_index_test.go @@ -0,0 +1,19 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestLatestAPIKeyIPIndexMigration(t *testing.T) { + content, err := FS.ReadFile("174_add_usage_logs_api_key_latest_ip_index_notx.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip") + require.Contains(t, sql, "ON usage_logs (api_key_id, created_at DESC, id DESC)") + require.Contains(t, sql, "INCLUDE (ip_address)") + require.Contains(t, sql, "WHERE ip_address IS NOT NULL AND ip_address <> ''") +} diff --git a/backend/migrations/openai_long_context_billing_migration_test.go b/backend/migrations/openai_long_context_billing_migration_test.go new file mode 100644 index 0000000000..212ac15d9d --- /dev/null +++ b/backend/migrations/openai_long_context_billing_migration_test.go @@ -0,0 +1,36 @@ +package migrations + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestMigration175DefaultsOrdinaryOpenAIAndInheritsForSparkShadows(t *testing.T) { + content, err := FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + + sql := string(content) + require.Contains(t, sql, "parent_account_id IS NULL") + require.Contains(t, sql, "quota_dimension = 'spark'") + require.Contains(t, sql, "parent.extra") + require.Contains(t, sql, "jsonb_typeof") + require.Contains(t, sql, "openai_long_context_billing_enabled") +} + +func TestMigration175GuardsMixedVersionAccountWrites(t *testing.T) { + content, err := FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + + sql := string(content) + require.Contains(t, sql, "RETURNS TRIGGER") + require.Contains(t, sql, "BEFORE INSERT OR UPDATE") + require.Contains(t, sql, "CREATE TRIGGER") + require.Contains(t, sql, "must be a boolean") + require.Contains(t, sql, "INSERT INTO scheduler_outbox") + require.Contains(t, sql, "'account_changed'") + require.Contains(t, sql, "jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean'") + require.Contains(t, sql, "WITH shadow_values AS") + require.Contains(t, sql, "TG_OP = 'UPDATE'") + require.Contains(t, sql, "OLD.extra->'openai_long_context_billing_enabled'") +} diff --git a/deploy/.env.example b/deploy/.env.example index 5925f0abb4..57056b4907 100644 --- a/deploy/.env.example +++ b/deploy/.env.example @@ -1,25 +1,37 @@ # ============================================================================= -# Sub2API Docker Environment Configuration +# Sub2API Container Environment Configuration # ============================================================================= # Copy this file to .env and modify as needed: # cp .env.example .env +# chmod 600 .env # nano .env # -# Then start with: docker-compose up -d +# Then start with Docker Compose or Apple container: +# docker compose up -d +# ./apple-container.sh up # ============================================================================= # ----------------------------------------------------------------------------- # Server Configuration # ----------------------------------------------------------------------------- -# Bind address for host port mapping +# IPv4 bind address for host port mapping BIND_HOST=0.0.0.0 -# Server port (exposed on host) +# Server port exposed on the host (Apple container requires 1025-65535) SERVER_PORT=8080 # Server mode: release or debug SERVER_MODE=release +# Return Server-Timing for authenticated requests made by the Admin web UI +ENABLE_SERVER_TIMING=false + +# Apple container image overrides (ignored by Docker Compose). Pin release tags +# or digests for repeatable operator-managed deployments. +APPLE_CONTAINER_SUB2API_IMAGE=weishaw/sub2api:latest +APPLE_CONTAINER_POSTGRES_IMAGE=postgres:18-alpine +APPLE_CONTAINER_REDIS_IMAGE=redis:8-alpine + # ----------------------------------------------------------------------------- # Logging Configuration # 日志配置 @@ -307,13 +319,15 @@ GATEWAY_SCHEDULING_OUTBOX_BACKLOG_REBUILD_ROWS=10000 GATEWAY_SCHEDULING_FULL_REBUILD_INTERVAL_SECONDS=300 # ----------------------------------------------------------------------------- -# Image Generation Stream & Concurrency (Optional) -# 图片生成流式与并发隔离配置(可选) +# Image Generation Keepalive & Concurrency (Optional) +# 图片生成保活与并发隔离配置(可选) # ----------------------------------------------------------------------------- # 图片流式上游数据间隔超时(秒)。0 表示禁用;非 0 时必须为 60-1800。 GATEWAY_IMAGE_STREAM_DATA_INTERVAL_TIMEOUT=900 # 图片流式 keepalive 间隔(秒)。0 表示禁用;非 0 时必须为 5-60。 GATEWAY_IMAGE_STREAM_KEEPALIVE_INTERVAL=10 +# 图片非流式 JSON keepalive 间隔(秒)。默认 0 禁用;首个心跳后 HTTP 状态会固化为 200。 +GATEWAY_IMAGE_NONSTREAM_KEEPALIVE_INTERVAL=0 # 是否启用进程级图片生成并发限制。默认 false,保持历史行为。 GATEWAY_IMAGE_CONCURRENCY_ENABLED=false # 当前进程允许同时处理的图片生成请求数。0 表示不限制。 diff --git a/deploy/APPLE_CONTAINER.md b/deploy/APPLE_CONTAINER.md new file mode 100644 index 0000000000..1133464a81 --- /dev/null +++ b/deploy/APPLE_CONTAINER.md @@ -0,0 +1,221 @@ +# Apple container Deployment + +Sub2API can run as a native three-service stack with Apple's `container` CLI. This workflow runs the published Sub2API, PostgreSQL, and Redis OCI images without Docker Desktop or a Docker-compatible daemon. + +## Support Level + +Apple `container` support is intended for local development and operator-managed deployments on a Mac. Docker Compose remains the recommended production deployment path. + +Apple `container` 1.1 does not provide restart policies, automatic startup, workload health scheduling, a Docker API socket, or full Compose orchestration. `apple-container.sh` supplies ordered startup and readiness checks when you invoke it, but it is not a continuously running supervisor. + +## Requirements + +- A Mac with Apple silicon +- macOS 26 or newer +- Apple `container` 1.1.0 or newer +- `openssl` for generating initial secrets +- Local Network access for `container-runtime-linux` when macOS prompts during the first published-container startup + +Install Apple `container` from its [official releases](https://github.com/apple/container/releases), then verify it: + +```bash +container --version +``` + +## Quick Start + +```bash +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api/deploy + +# Creates .env with random PostgreSQL, JWT, and TOTP secrets. +./apple-container.sh init + +# Review optional settings before startup. +nano .env + +# Creates volumes/network/containers, waits for dependencies, and starts Sub2API. +./apple-container.sh up + +# Verifies PostgreSQL, Redis, and the application endpoint. +./apple-container.sh status +``` + +Open `http://localhost:8080`. If `ADMIN_PASSWORD` is empty, retrieve the generated password with: + +```bash +./apple-container.sh logs app +``` + +The env file uses literal `KEY=value` syntax. Do not use Compose expressions such as `${VALUE:-default}`, and do not quote values unless the quote characters are part of the intended value. `BIND_HOST` must be an IPv4 address, and `SERVER_PORT` must be between 1025 and 65535. + +## Commands + +```bash +# Start dependencies and recreate the lightweight app container with current IPs. +./apple-container.sh up + +# Also recreate PostgreSQL and Redis containers, preserving their volumes. +./apple-container.sh up --recreate + +# Stop containers while preserving all resources and data. +./apple-container.sh down + +# Restart PostgreSQL, Redis, and Sub2API in dependency order. +./apple-container.sh restart + +# Show resource state and run live health probes. +./apple-container.sh status + +# Follow one service's logs. +./apple-container.sh logs app -f +./apple-container.sh logs postgres -f +./apple-container.sh logs redis -f + +# Pull all configured images for linux/arm64, then recreate containers. +./apple-container.sh pull +./apple-container.sh up --recreate + +# Delete containers and the network, preserving named volumes. +./apple-container.sh destroy --yes + +# Permanently delete the stack and all application/database/cache data. +./apple-container.sh destroy --volumes --yes +``` + +`destroy --volumes` does not remove `.env`, backup files, or pulled images. Delete credentials and backups separately when decommissioning a deployment. Use `container image delete ` only after confirming no other Apple containers use that image. + +After a host reboot or `container system stop`, run `./apple-container.sh up` again. Apple `container` does not automatically restart persisted containers. + +## Configuration + +The script uses `deploy/.env`, the same source file used by Docker Compose. Export `SUB2API_ENV_FILE` to use another file for every command in the current shell: + +```bash +export SUB2API_ENV_FILE=/absolute/path/to/sub2api.env +./apple-container.sh init +./apple-container.sh up +``` + +Apple-specific image overrides are available: + +```dotenv +APPLE_CONTAINER_SUB2API_IMAGE=weishaw/sub2api:latest +APPLE_CONTAINER_POSTGRES_IMAGE=postgres:18-alpine +APPLE_CONTAINER_REDIS_IMAGE=redis:8-alpine +``` + +The normal `up` command recreates the application container, so application environment changes are applied immediately. Use `up --recreate` when changing PostgreSQL or Redis container images or Redis runtime configuration. Persistent data remains in named volumes. + +`POSTGRES_USER`, `POSTGRES_PASSWORD`, and `POSTGRES_DB` are applied only when PostgreSQL initializes an empty data volume. Changing them in `.env` and recreating the container does not change an existing database. Rotate a password with `ALTER ROLE`, and plan explicit migrations for user or database changes. To intentionally initialize a new empty database, first back up the old one and use `destroy --volumes`. + +Apple-specific handling of shared settings: + +| Setting | Apple workflow behavior | +|---|---| +| Application and gateway variables | Passed to Sub2API from `.env` | +| `BIND_HOST`, `SERVER_PORT` | Used for the macOS published port | +| `POSTGRES_USER`, `POSTGRES_PASSWORD`, `POSTGRES_DB` | PostgreSQL first initialization only | +| `REDIS_PASSWORD` | Applied to Redis and Sub2API | +| `DATABASE_PORT`, `REDIS_PORT` | Internal ports are fixed to 5432 and 6379 | +| `POSTGRES_MAX_*`, `REDIS_MAXCLIENTS` | Not currently applied to the database/cache server | + +## Managed Resources + +The script creates only resources carrying the `org.sub2api.stack=apple-container` label: + +| Type | Names | +|---|---| +| Containers | `sub2api-apple`, `sub2api-apple-postgres`, `sub2api-apple-redis` | +| Network | `sub2api-apple` | +| Volumes | `sub2api-apple-data`, `sub2api-apple-postgres-data`, `sub2api-apple-redis-data` | + +The PostgreSQL volume is mounted at `/var/lib/postgresql`, retaining PostgreSQL 18's default child data directory. Sub2API and Redis also store data in child directories below their Apple volume mount points. This is required because Apple named volumes do not have Docker's copy-up and mount-point ownership behavior. + +## Networking + +Apple `container` 1.1 does not provide Compose-style network-scoped service aliases. After PostgreSQL and Redis start, the script reads their current private-network IPv4 addresses from `container inspect`, injects those addresses into a newly created application container, and then starts Sub2API. The script does not modify `~/.config/container/config.toml` or the macOS host resolver. + +All three services attach only to the private `sub2api-apple` network. Only the application publishes a host port; database and Redis ports remain unpublished. + +The application container is intentionally recreated by every `up` and `restart` operation because dependency VM addresses can change after they stop. Application data remains in `sub2api-apple-data`. + +The script checks the published `/health` endpoint from macOS before reporting success. Approve the Local Network prompt on first startup. If the internal probe succeeds but the host-port probe fails with a connection reset, enable Local Network access for `container-runtime-linux`, run `container system stop` followed by `container system start`, and then run `up` again. Runtime upgrades may prompt for permission again. + +## Backup and Upgrade + +Pin image release tags or digests in `.env` before using this workflow for persistent data. Before an application or database image upgrade, create backups while the stack is healthy: + +```bash +umask 077 +mkdir -p backups + +# Logical PostgreSQL backup. +container exec sub2api-apple sh -c \ + 'PGPASSWORD="$DATABASE_PASSWORD" pg_dump -h "$DATABASE_HOST" -U "$DATABASE_USER" "$DATABASE_DBNAME"' \ + > backups/sub2api.sql + +# Application configuration and local files. +container exec sub2api-apple sh -c 'tar -C "$DATA_DIR" -czf - .' \ + > backups/sub2api-data.tar.gz + +./apple-container.sh pull +./apple-container.sh up --recreate +./apple-container.sh status +``` + +Database migrations are forward-only. Keep the previous image reference and both backups until the upgraded stack has been validated; image rollback alone cannot reverse a migrated database. Test restore procedures before relying on this workflow for important data. + +To restore these backups into an existing stack, first ensure the image versions are compatible with the backup, then stop writers and replace both data sets: + +```bash +# Ensure empty/current resources exist, then stop the stack. +./apple-container.sh up +./apple-container.sh down + +# Remove only the app container so a helper can mount its named volume. +container delete sub2api-apple +SUB2API_IMAGE=weishaw/sub2api:latest # Match APPLE_CONTAINER_SUB2API_IMAGE in .env. +container run --rm --name sub2api-apple-data-restore \ + --entrypoint /bin/sh \ + --volume sub2api-apple-data:/restore \ + --volume "$PWD/backups:/backup:ro" \ + "$SUB2API_IMAGE" \ + -c 'rm -rf /restore/data && mkdir -p /restore/data && tar -xzf /backup/sub2api-data.tar.gz -C /restore/data' + +# Restore the logical database while the application is absent. +container start sub2api-apple-postgres +until container exec sub2api-apple-postgres sh -c 'pg_isready -U "$POSTGRES_USER" -d "$POSTGRES_DB"'; do sleep 1; done +container copy backups/sub2api.sql sub2api-apple-postgres:/tmp/sub2api.sql +container exec sub2api-apple-postgres sh -c ' + export PGPASSWORD="$POSTGRES_PASSWORD" + dropdb -h 127.0.0.1 -U "$POSTGRES_USER" --if-exists --force "$POSTGRES_DB" + createdb -h 127.0.0.1 -U "$POSTGRES_USER" "$POSTGRES_DB" + psql -h 127.0.0.1 -U "$POSTGRES_USER" -d "$POSTGRES_DB" -v ON_ERROR_STOP=1 -f /tmp/sub2api.sql + rm /tmp/sub2api.sql +' + +./apple-container.sh up +./apple-container.sh status +``` + +For disaster recovery after deleting the named volumes, run `up` once to create a fresh stack before following the restore sequence. Perform restore drills with non-production data first. + +To upgrade the Apple runtime itself: + +```bash +./apple-container.sh down +container system stop +# Install/update Apple container 1.1.0 or newer. +container system start +./apple-container.sh up +``` + +## Operational Limitations + +- There is no `restart: unless-stopped` equivalent. Run `up` after reboot, or add your own launchd supervisor. +- Health probes run during `up`, `restart`, and `status`; Apple `container` does not continuously schedule them. +- Docker Compose, Testcontainers, Buildx, and tools requiring `/var/run/docker.sock` cannot use this runtime directly. +- Named volume backup and restore must be tested before using this workflow for important data. +- The script targets native `linux/arm64` images. The normal Sub2API release publishes an arm64 variant. +- Runtime environment values, including credentials, are retained in Apple container configuration and are visible to users who can inspect the local runtime. diff --git a/deploy/README.md b/deploy/README.md index dd311721d9..d4fcc133b2 100644 --- a/deploy/README.md +++ b/deploy/README.md @@ -1,12 +1,13 @@ # Sub2API Deployment Files -This directory contains files for deploying Sub2API on Linux servers. +This directory contains files for deploying Sub2API on Linux servers and Apple-silicon Macs. ## Deployment Methods | Method | Best For | Setup Wizard | |--------|----------|--------------| | **Docker Compose** | Quick setup, all-in-one | Not needed (auto-setup) | +| **Apple container** | Native local stack on macOS 26 | Not needed (auto-setup) | | **Binary Install** | Production servers, systemd | Web-based wizard | ## Files @@ -16,7 +17,9 @@ This directory contains files for deploying Sub2API on Linux servers. | `docker-compose.yml` | Docker Compose configuration (named volumes) | | `docker-compose.local.yml` | Docker Compose configuration (local directories, easy migration) | | `docker-deploy.sh` | **One-click Docker deployment script (recommended)** | -| `.env.example` | Docker environment variables template | +| `apple-container.sh` | Native Apple `container` lifecycle script | +| `APPLE_CONTAINER.md` | Apple `container` deployment and operations guide | +| `.env.example` | Container environment variables template | | `DOCKER.md` | Docker Hub documentation | | `install.sh` | One-click binary installation script | | `install-datamanagementd.sh` | datamanagementd 一键安装脚本 | @@ -27,6 +30,23 @@ This directory contains files for deploying Sub2API on Linux servers. --- +## Apple container Deployment + +Apple-silicon Macs running macOS 26 can run the complete Sub2API, PostgreSQL, and Redis stack with Apple `container` 1.1.0 or newer: + +```bash +./apple-container.sh init +./apple-container.sh up +./apple-container.sh status +./apple-container.sh logs app -f +``` + +The script uses Apple named volumes, starts dependencies in order, and performs live readiness checks. It does not provide a continuous restart supervisor; run `./apple-container.sh up` after a host reboot. Docker Compose remains the recommended production deployment path. + +See [APPLE_CONTAINER.md](./APPLE_CONTAINER.md) for configuration, upgrades, persistence, networking behavior, and limitations. + +--- + ## Docker Deployment (Recommended) ### Method 1: One-Click Deployment (Recommended) @@ -76,6 +96,7 @@ cd sub2api/deploy # Configure environment cp .env.example .env +chmod 600 .env nano .env # Set POSTGRES_PASSWORD and other required variables # Generate secure secrets (recommended) diff --git a/deploy/apple-container.sh b/deploy/apple-container.sh new file mode 100755 index 0000000000..5de4d5cab1 --- /dev/null +++ b/deploy/apple-container.sh @@ -0,0 +1,926 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +ENV_FILE="${SUB2API_ENV_FILE:-${SCRIPT_DIR}/.env}" + +STACK_LABEL_KEY="org.sub2api.stack" +STACK_LABEL_VALUE="apple-container" +NETWORK_NAME="sub2api-apple" +APP_CONTAINER="sub2api-apple" +POSTGRES_CONTAINER="sub2api-apple-postgres" +REDIS_CONTAINER="sub2api-apple-redis" +APP_VOLUME="sub2api-apple-data" +POSTGRES_VOLUME="sub2api-apple-postgres-data" +REDIS_VOLUME="sub2api-apple-redis-data" +PLATFORM="linux/arm64" + +TEMP_DIR="" +LOCK_DIR="${TMPDIR:-/tmp}/sub2api-apple-container.lock" +LOCK_ACQUIRED=false + +APP_IMAGE="" +POSTGRES_IMAGE="" +REDIS_IMAGE="" +BIND_HOST="" +HOST_PORT="" +ACCESS_HOST="" +POSTGRES_USER="" +POSTGRES_PASSWORD="" +POSTGRES_DB="" +REDIS_PASSWORD="" +TZ_VALUE="" +POSTGRES_ADDRESS="" +REDIS_ADDRESS="" +APP_ENV_FILE="" +POSTGRES_ENV_FILE="" +POSTGRES_PROBE_ENV_FILE="" +REDIS_ENV_FILE="" + +info() { + printf '[INFO] %s\n' "$*" +} + +warn() { + printf '[WARN] %s\n' "$*" >&2 +} + +die() { + printf '[ERROR] %s\n' "$*" >&2 + exit 1 +} + +usage() { + cat <<'EOF' +Usage: ./apple-container.sh [options] + +Commands: + init Create .env and generate required secrets + up [--recreate] Create and start the complete Sub2API stack + down Stop the stack and preserve all data + restart Restart the stack in dependency order + status Show container and workload health + logs [-f] Show logs for app, postgres, or redis + pull Pull all stack images for linux/arm64 + destroy [options] Delete stack containers and network + +Destroy options: + --volumes Also delete all persistent data volumes + --yes Skip the confirmation prompt + +Environment: + SUB2API_ENV_FILE Path to the deployment env file (default: deploy/.env) +EOF +} + +cleanup() { + local exit_code=$? + + if [[ -n "${TEMP_DIR}" && -d "${TEMP_DIR}" ]]; then + rm -rf "${TEMP_DIR}" + fi + if [[ "${LOCK_ACQUIRED}" == true && -d "${LOCK_DIR}" ]]; then + rm -f "${LOCK_DIR}/pid" + rmdir "${LOCK_DIR}" 2>/dev/null || true + fi + + exit "${exit_code}" +} + +acquire_lock() { + if ! mkdir "${LOCK_DIR}" 2>/dev/null; then + local owner_pid="" + if [[ -f "${LOCK_DIR}/pid" ]]; then + owner_pid="$(<"${LOCK_DIR}/pid")" + fi + if [[ "${owner_pid}" =~ ^[0-9]+$ ]] && ! kill -0 "${owner_pid}" 2>/dev/null; then + rm -rf "${LOCK_DIR}" + mkdir "${LOCK_DIR}" || die "Failed to reclaim stale operation lock." + else + die "Another Sub2API Apple container operation is already running." + fi + fi + printf '%s\n' "$$" >"${LOCK_DIR}/pid" + LOCK_ACQUIRED=true + trap cleanup EXIT + trap 'exit 130' INT + trap 'exit 143' TERM + trap 'exit 129' HUP +} + +require_command() { + command -v "$1" >/dev/null 2>&1 || die "Required command not found: $1" +} + +require_container_version() { + local version_output major minor + + require_command container + require_command plutil + version_output="$(container --version)" + if [[ ! "${version_output}" =~ ([0-9]+)\.([0-9]+)\.([0-9]+) ]]; then + die "Unable to parse Apple container version: ${version_output}" + fi + + major="${BASH_REMATCH[1]}" + minor="${BASH_REMATCH[2]}" + if (( major < 1 || (major == 1 && minor < 1) )); then + die "Apple container 1.1.0 or newer is required; found ${version_output}." + fi +} + +system_is_running() { + container system status >/dev/null 2>&1 +} + +start_system() { + if ! system_is_running; then + info "Starting Apple container services..." + container system start --enable-kernel-install + fi +} + +list_resource_ids() { + case "$1" in + container) container list --all --quiet ;; + network) container network list --quiet ;; + volume) container volume list --quiet ;; + *) die "Unknown resource type: $1" ;; + esac +} + +resource_exists() { + local resource_type=$1 + local resource_name=$2 + local output line + + if ! output="$(list_resource_ids "${resource_type}")"; then + die "Failed to list Apple container ${resource_type} resources." + fi + + while IFS= read -r line; do + if [[ "${line}" == "${resource_name}" ]]; then + return 0 + fi + done <<<"${output}" + + return 1 +} + +inspect_resource() { + case "$1" in + container) container inspect "$2" ;; + network) container network inspect "$2" ;; + volume) container volume inspect "$2" ;; + *) die "Unknown resource type: $1" ;; + esac +} + +assert_resource_owned() { + local resource_type=$1 + local resource_name=$2 + local inspection compact + + inspection="$(inspect_resource "${resource_type}" "${resource_name}" | \ + plutil -extract 0.configuration.labels json -o - -)" || \ + die "Failed to inspect ${resource_type} ${resource_name}." + compact="$(printf '%s' "${inspection}" | tr -d '[:space:]')" + if [[ "${compact}" != *"\"${STACK_LABEL_KEY}\":\"${STACK_LABEL_VALUE}\""* ]]; then + die "Refusing to manage existing ${resource_type} '${resource_name}' because it is not owned by this stack." + fi +} + +preflight_stack_ownership() { + local resource_name + + for resource_name in "${APP_CONTAINER}" "${REDIS_CONTAINER}" "${POSTGRES_CONTAINER}"; do + if resource_exists container "${resource_name}"; then + assert_resource_owned container "${resource_name}" + fi + done + if resource_exists network "${NETWORK_NAME}"; then + assert_resource_owned network "${NETWORK_NAME}" + fi + for resource_name in "${APP_VOLUME}" "${REDIS_VOLUME}" "${POSTGRES_VOLUME}"; do + if resource_exists volume "${resource_name}"; then + assert_resource_owned volume "${resource_name}" + fi + done +} + +ensure_network() { + if resource_exists network "${NETWORK_NAME}"; then + assert_resource_owned network "${NETWORK_NAME}" + return + fi + + info "Creating network ${NETWORK_NAME}..." + container network create \ + --label "${STACK_LABEL_KEY}=${STACK_LABEL_VALUE}" \ + "${NETWORK_NAME}" >/dev/null +} + +ensure_volume() { + local volume_name=$1 + + if resource_exists volume "${volume_name}"; then + assert_resource_owned volume "${volume_name}" + return + fi + + info "Creating volume ${volume_name}..." + container volume create \ + --label "${STACK_LABEL_KEY}=${STACK_LABEL_VALUE}" \ + "${volume_name}" >/dev/null +} + +ensure_image_available() { + local image=$1 + + if container image inspect "${image}" >/dev/null 2>&1; then + return + fi + info "Pulling ${image}..." + container image pull --platform "${PLATFORM}" "${image}" +} + +container_is_running() { + local container_name=$1 + local output line + + output="$(container list --quiet)" || die "Failed to list running Apple containers." + while IFS= read -r line; do + if [[ "${line}" == "${container_name}" ]]; then + return 0 + fi + done <<<"${output}" + + return 1 +} + +ensure_system() { + require_container_version + require_command curl + start_system +} + +container_ipv4_address() { + local container_name=$1 + local address + + address="$(container inspect "${container_name}" | \ + plutil -extract 0.status.networks.0.ipv4Address raw -o - -)" || \ + die "Unable to read the network address for ${container_name}." + address="${address%%/*}" + [[ "${address}" =~ ^[0-9]+\.[0-9]+\.[0-9]+\.[0-9]+$ ]] || \ + die "Apple container returned an invalid IPv4 address for ${container_name}: ${address}" + printf '%s\n' "${address}" +} + +read_env_value() { + local key=$1 + local fallback=${2-} + + awk -v wanted="${key}" -v fallback="${fallback}" ' + BEGIN { found = 0 } + /^[[:space:]]*#/ || /^[[:space:]]*$/ { next } + { + separator = index($0, "=") + if (separator == 0) { next } + key = substr($0, 1, separator - 1) + if (key == wanted) { + value = substr($0, separator + 1) + sub(/\r$/, "", value) + found = 1 + } + } + END { + if (found) { print value } + else { print fallback } + } + ' "${ENV_FILE}" +} + +replace_env_value() { + local key=$1 + local value=$2 + local target_file=${3:-${ENV_FILE}} + local temp_file="${target_file}.tmp.$$" + + awk -v wanted="${key}" -v replacement="${value}" ' + BEGIN { replaced = 0 } + { + separator = index($0, "=") + key = separator == 0 ? "" : substr($0, 1, separator - 1) + if (key == wanted) { + if (!replaced) { print wanted "=" replacement } + replaced = 1 + next + } + print + } + END { + if (!replaced) { print wanted "=" replacement } + } + ' "${target_file}" >"${temp_file}" + chmod 600 "${temp_file}" + mv "${temp_file}" "${target_file}" +} + +generate_secret() { + openssl rand -hex 32 +} + +cmd_init() { + local env_dir temp_file postgres_secret jwt_secret totp_secret + + require_command openssl + + if [[ -e "${ENV_FILE}" ]]; then + die "Environment file already exists: ${ENV_FILE}" + fi + + postgres_secret="$(generate_secret)" || die "Failed to generate PostgreSQL password." + jwt_secret="$(generate_secret)" || die "Failed to generate JWT secret." + totp_secret="$(generate_secret)" || die "Failed to generate TOTP encryption key." + [[ -n "${postgres_secret}" && -n "${jwt_secret}" && -n "${totp_secret}" ]] || \ + die "Secret generation returned an empty value." + + env_dir="$(dirname "${ENV_FILE}")" + temp_file="${ENV_FILE}.init.tmp.$$" + mkdir -p "${env_dir}" + cp "${SCRIPT_DIR}/.env.example" "${temp_file}" + chmod 600 "${temp_file}" + replace_env_value POSTGRES_PASSWORD "${postgres_secret}" "${temp_file}" + replace_env_value JWT_SECRET "${jwt_secret}" "${temp_file}" + replace_env_value TOTP_ENCRYPTION_KEY "${totp_secret}" "${temp_file}" + mv "${temp_file}" "${ENV_FILE}" + + info "Created ${ENV_FILE} with generated secrets." + info "Review the file, then run: SUB2API_ENV_FILE='${ENV_FILE}' ${SCRIPT_DIR}/apple-container.sh up" +} + +validate_port() { + local port=$1 + local decimal_port + + [[ "${port}" =~ ^[0-9]+$ ]] || die "SERVER_PORT must be numeric: ${port}" + decimal_port=$((10#${port})) + (( decimal_port >= 1025 && decimal_port <= 65535 )) || \ + die "SERVER_PORT must be between 1025 and 65535 for Apple container port forwarding." +} + +validate_ipv4_address() { + local address=$1 + local first second third fourth extra octet + + IFS=. read -r first second third fourth extra <<<"${address}" + [[ -n "${first}" && -n "${second}" && -n "${third}" && -n "${fourth}" && -z "${extra}" ]] || \ + die "BIND_HOST must be a valid IPv4 address: ${address}" + for octet in "${first}" "${second}" "${third}" "${fourth}"; do + [[ "${octet}" =~ ^[0-9]+$ ]] || die "BIND_HOST must be a valid IPv4 address: ${address}" + (( 10#${octet} <= 255 )) || die "BIND_HOST must be a valid IPv4 address: ${address}" + done +} + +validate_env_file_security() { + local owner mode permissions + + [[ -f "${ENV_FILE}" ]] || die "Environment file not found: ${ENV_FILE}. Run '$0 init' first." + owner="$(stat -f '%u' "${ENV_FILE}")" || die "Unable to read owner for ${ENV_FILE}." + mode="$(stat -f '%Lp' "${ENV_FILE}")" || die "Unable to read permissions for ${ENV_FILE}." + [[ "${owner}" == "${EUID}" ]] || die "Environment file must be owned by the current user: ${ENV_FILE}" + [[ "${mode}" =~ ^[0-7]+$ ]] || die "Unable to parse permissions for ${ENV_FILE}: ${mode}" + permissions=$((8#${mode})) + (( (permissions & 077) == 0 )) || \ + die "Environment file must not be readable by group or others. Run: chmod 600 '${ENV_FILE}'" +} + +prepare_environment() { + validate_env_file_security + + APP_IMAGE="$(read_env_value APPLE_CONTAINER_SUB2API_IMAGE weishaw/sub2api:latest)" + POSTGRES_IMAGE="$(read_env_value APPLE_CONTAINER_POSTGRES_IMAGE postgres:18-alpine)" + REDIS_IMAGE="$(read_env_value APPLE_CONTAINER_REDIS_IMAGE redis:8-alpine)" + BIND_HOST="$(read_env_value BIND_HOST 0.0.0.0)" + HOST_PORT="$(read_env_value SERVER_PORT 8080)" + POSTGRES_USER="$(read_env_value POSTGRES_USER sub2api)" + POSTGRES_PASSWORD="$(read_env_value POSTGRES_PASSWORD)" + POSTGRES_DB="$(read_env_value POSTGRES_DB sub2api)" + REDIS_PASSWORD="$(read_env_value REDIS_PASSWORD)" + TZ_VALUE="$(read_env_value TZ Asia/Shanghai)" + + [[ -n "${BIND_HOST}" ]] || die "BIND_HOST must not be empty." + validate_ipv4_address "${BIND_HOST}" + validate_port "${HOST_PORT}" + if [[ "${BIND_HOST}" == "0.0.0.0" ]]; then + ACCESS_HOST="127.0.0.1" + else + ACCESS_HOST="${BIND_HOST}" + fi + [[ -n "${POSTGRES_USER}" ]] || die "POSTGRES_USER must not be empty." + [[ -n "${POSTGRES_DB}" ]] || die "POSTGRES_DB must not be empty." + if [[ -z "${POSTGRES_PASSWORD}" || "${POSTGRES_PASSWORD}" == "change_this_secure_password" ]]; then + die "Set a secure POSTGRES_PASSWORD in ${ENV_FILE}." + fi + + TEMP_DIR="$(mktemp -d "${TMPDIR:-/tmp}/sub2api-apple.XXXXXX")" + APP_ENV_FILE="${TEMP_DIR}/app.env" + POSTGRES_ENV_FILE="${TEMP_DIR}/postgres.env" + POSTGRES_PROBE_ENV_FILE="${TEMP_DIR}/postgres-probe.env" + REDIS_ENV_FILE="${TEMP_DIR}/redis.env" + + cat >"${POSTGRES_ENV_FILE}" <"${POSTGRES_PROBE_ENV_FILE}" <"${REDIS_ENV_FILE}" <>"${REDIS_ENV_FILE}" + fi + + chmod 600 "${POSTGRES_ENV_FILE}" "${POSTGRES_PROBE_ENV_FILE}" "${REDIS_ENV_FILE}" +} + +prepare_app_environment() { + [[ -n "${POSTGRES_ADDRESS}" && -n "${REDIS_ADDRESS}" ]] || \ + die "Dependency network addresses are not available." + + cp "${ENV_FILE}" "${APP_ENV_FILE}" + cat >>"${APP_ENV_FILE}" </dev/null +} + +create_redis_container() { + info "Creating Redis container..." + container create \ + --name "${REDIS_CONTAINER}" \ + --label "${STACK_LABEL_KEY}=${STACK_LABEL_VALUE}" \ + --network "${NETWORK_NAME}" \ + --platform "${PLATFORM}" \ + --ulimit nofile=100000:100000 \ + --env-file "${REDIS_ENV_FILE}" \ + --volume "${REDIS_VOLUME}:/var/lib/redis" \ + "${REDIS_IMAGE}" \ + sh -c 'set -e; mkdir -p /var/lib/redis/data; chown redis:redis /var/lib/redis/data; exec /usr/local/bin/docker-entrypoint.sh redis-server --dir /var/lib/redis/data --save 60 1 --appendonly yes --appendfsync everysec ${REDIS_PASSWORD:+--requirepass "$REDIS_PASSWORD"}' \ + >/dev/null +} + +create_app_container() { + info "Creating Sub2API container..." + container create \ + --name "${APP_CONTAINER}" \ + --label "${STACK_LABEL_KEY}=${STACK_LABEL_VALUE}" \ + --network "${NETWORK_NAME}" \ + --platform "${PLATFORM}" \ + --ulimit nofile=100000:100000 \ + --publish "${BIND_HOST}:${HOST_PORT}:8080/tcp" \ + --env-file "${APP_ENV_FILE}" \ + --volume "${APP_VOLUME}:/app/storage" \ + --entrypoint /bin/sh \ + "${APP_IMAGE}" \ + -c 'set -e; mkdir -p "$DATA_DIR"; chown -R sub2api:sub2api "$DATA_DIR"; exec su-exec sub2api /app/sub2api' \ + >/dev/null +} + +ensure_container() { + local container_name=$1 + local create_function=$2 + + if resource_exists container "${container_name}"; then + assert_resource_owned container "${container_name}" + return + fi + + "${create_function}" +} + +start_container_if_needed() { + local container_name=$1 + + if container_is_running "${container_name}"; then + return + fi + + info "Starting ${container_name}..." + container start "${container_name}" >/dev/null +} + +stop_container_if_running() { + local container_name=$1 + + if ! resource_exists container "${container_name}"; then + return + fi + assert_resource_owned container "${container_name}" + if container_is_running "${container_name}"; then + info "Stopping ${container_name}..." + container stop --time 30 "${container_name}" >/dev/null + fi +} + +delete_container_if_present() { + local container_name=$1 + + if ! resource_exists container "${container_name}"; then + return + fi + assert_resource_owned container "${container_name}" + if container_is_running "${container_name}"; then + container stop --time 30 "${container_name}" >/dev/null + fi + info "Deleting ${container_name}..." + container delete "${container_name}" >/dev/null +} + +wait_for_probe() { + local description=$1 + local attempts=$2 + shift 2 + + local attempt + for ((attempt = 1; attempt <= attempts; attempt++)); do + if "$@" >/dev/null 2>&1; then + info "${description} is ready." + return 0 + fi + sleep 1 + done + + return 1 +} + +probe_postgres() { + container exec --env-file "${POSTGRES_PROBE_ENV_FILE}" \ + "${POSTGRES_CONTAINER}" \ + psql -h 127.0.0.1 -U "${POSTGRES_USER}" -d "${POSTGRES_DB}" \ + -v ON_ERROR_STOP=1 -tAc 'SELECT 1' +} + +probe_redis() { + container exec --env-file "${REDIS_ENV_FILE}" \ + "${REDIS_CONTAINER}" \ + redis-cli ping +} + +probe_app() { + container exec "${APP_CONTAINER}" \ + wget -q -T 5 -O /dev/null http://localhost:8080/health +} + +probe_host_app() { + curl --fail --silent --show-error --max-time 5 \ + "http://${ACCESS_HOST}:${HOST_PORT}/health" +} + +show_failure_logs() { + local container_name=$1 + + warn "Last logs from ${container_name}:" + container logs -n 50 "${container_name}" >&2 || true +} + +start_dependencies() { + start_container_if_needed "${POSTGRES_CONTAINER}" + if ! wait_for_probe "PostgreSQL" 90 probe_postgres; then + show_failure_logs "${POSTGRES_CONTAINER}" + die "PostgreSQL did not become ready." + fi + + start_container_if_needed "${REDIS_CONTAINER}" + if ! wait_for_probe "Redis" 60 probe_redis; then + show_failure_logs "${REDIS_CONTAINER}" + die "Redis did not become ready." + fi +} + +start_app() { + start_container_if_needed "${APP_CONTAINER}" + if ! wait_for_probe "Sub2API" 180 probe_app; then + show_failure_logs "${APP_CONTAINER}" + die "Sub2API did not become ready." + fi + if ! wait_for_probe "Sub2API host port" 15 probe_host_app; then + die "Host port forwarding failed. In System Settings > Privacy & Security > Local Network, allow container-runtime-linux; restart Apple container services; then run 'apple-container.sh up' again." + fi +} + +cmd_up() { + local recreate=false + + if [[ $# -gt 1 || ($# -eq 1 && "${1-}" != "--recreate") ]]; then + usage + exit 2 + fi + if [[ $# -eq 1 ]]; then + recreate=true + fi + + ensure_system + prepare_environment + preflight_stack_ownership + ensure_network + ensure_volume "${APP_VOLUME}" + ensure_volume "${POSTGRES_VOLUME}" + ensure_volume "${REDIS_VOLUME}" + ensure_image_available "${APP_IMAGE}" + ensure_image_available "${POSTGRES_IMAGE}" + ensure_image_available "${REDIS_IMAGE}" + + if [[ "${recreate}" == true ]]; then + delete_container_if_present "${APP_CONTAINER}" + delete_container_if_present "${REDIS_CONTAINER}" + delete_container_if_present "${POSTGRES_CONTAINER}" + fi + + ensure_container "${POSTGRES_CONTAINER}" create_postgres_container + ensure_container "${REDIS_CONTAINER}" create_redis_container + start_dependencies + POSTGRES_ADDRESS="$(container_ipv4_address "${POSTGRES_CONTAINER}")" + REDIS_ADDRESS="$(container_ipv4_address "${REDIS_CONTAINER}")" + prepare_app_environment + # The dependency IPs may change whenever their lightweight VMs restart. + delete_container_if_present "${APP_CONTAINER}" + create_app_container + start_app + + info "Sub2API is available at http://${ACCESS_HOST}:${HOST_PORT}" +} + +cmd_down() { + require_container_version + if ! system_is_running; then + info "Apple container services are already stopped." + return + fi + preflight_stack_ownership + stop_container_if_running "${APP_CONTAINER}" + stop_container_if_running "${REDIS_CONTAINER}" + stop_container_if_running "${POSTGRES_CONTAINER}" + info "Sub2API stack stopped; persistent volumes were preserved." +} + +cmd_restart() { + cmd_down + cmd_up +} + +print_container_status() { + local service=$1 + local container_name=$2 + + if ! resource_exists container "${container_name}"; then + printf '%-12s %s\n' "${service}" "missing" + elif container_is_running "${container_name}"; then + printf '%-12s %s\n' "${service}" "running" + else + printf '%-12s %s\n' "${service}" "stopped" + fi +} + +cmd_status() { + local failed=0 + + require_container_version + if ! system_is_running; then + printf '%-12s %s\n' "system" "stopped" + return 1 + fi + + printf '%-12s %s\n' "system" "running" + preflight_stack_ownership + print_container_status app "${APP_CONTAINER}" + print_container_status postgres "${POSTGRES_CONTAINER}" + print_container_status redis "${REDIS_CONTAINER}" + + if [[ -f "${ENV_FILE}" ]]; then + prepare_environment + if container_is_running "${POSTGRES_CONTAINER}" && probe_postgres >/dev/null 2>&1; then + printf '%-12s %s\n' "postgres" "healthy" + else + printf '%-12s %s\n' "postgres" "unhealthy" + failed=1 + fi + if container_is_running "${REDIS_CONTAINER}" && probe_redis >/dev/null 2>&1; then + printf '%-12s %s\n' "redis" "healthy" + else + printf '%-12s %s\n' "redis" "unhealthy" + failed=1 + fi + if container_is_running "${APP_CONTAINER}" && probe_app >/dev/null 2>&1; then + printf '%-12s %s\n' "app" "healthy" + else + printf '%-12s %s\n' "app" "unhealthy" + failed=1 + fi + if container_is_running "${APP_CONTAINER}" && probe_host_app >/dev/null 2>&1; then + printf '%-12s %s\n' "host-port" "healthy" + else + printf '%-12s %s\n' "host-port" "unhealthy" + failed=1 + fi + else + warn "Health probes require ${ENV_FILE}." + failed=1 + fi + + return "${failed}" +} + +cmd_logs() { + local service=${1-} + local follow=${2-} + local container_name + + [[ $# -ge 1 && $# -le 2 ]] || { usage; exit 2; } + if [[ -n "${follow}" && "${follow}" != "-f" && "${follow}" != "--follow" ]]; then + usage + exit 2 + fi + + case "${service}" in + app|sub2api) container_name="${APP_CONTAINER}" ;; + postgres) container_name="${POSTGRES_CONTAINER}" ;; + redis) container_name="${REDIS_CONTAINER}" ;; + *) die "Unknown service '${service}'. Use app, postgres, or redis." ;; + esac + + require_container_version + system_is_running || die "Apple container services are not running." + resource_exists container "${container_name}" || die "Container not found: ${container_name}" + assert_resource_owned container "${container_name}" + if [[ -n "${follow}" ]]; then + container logs --follow "${container_name}" + else + container logs "${container_name}" + fi +} + +cmd_pull() { + ensure_system + prepare_environment + info "Pulling ${APP_IMAGE}..." + container image pull --platform "${PLATFORM}" "${APP_IMAGE}" + info "Pulling ${POSTGRES_IMAGE}..." + container image pull --platform "${PLATFORM}" "${POSTGRES_IMAGE}" + info "Pulling ${REDIS_IMAGE}..." + container image pull --platform "${PLATFORM}" "${REDIS_IMAGE}" +} + +confirm_destroy() { + local include_volumes=$1 + local answer + + if [[ "${include_volumes}" == true ]]; then + printf 'Delete the Sub2API stack and all persistent data? [y/N] ' + else + printf 'Delete the Sub2API containers and network, preserving volumes? [y/N] ' + fi + read -r answer + [[ "${answer}" == "y" || "${answer}" == "Y" ]] +} + +delete_volume_if_present() { + local volume_name=$1 + + if resource_exists volume "${volume_name}"; then + assert_resource_owned volume "${volume_name}" + info "Deleting volume ${volume_name}..." + container volume delete "${volume_name}" >/dev/null + fi +} + +cmd_destroy() { + local include_volumes=false + local assume_yes=false + local argument + + for argument in "$@"; do + case "${argument}" in + --volumes) include_volumes=true ;; + --yes) assume_yes=true ;; + *) usage; exit 2 ;; + esac + done + + require_container_version + start_system + preflight_stack_ownership + if [[ "${assume_yes}" != true ]] && ! confirm_destroy "${include_volumes}"; then + info "Cancelled." + return + fi + + delete_container_if_present "${APP_CONTAINER}" + delete_container_if_present "${REDIS_CONTAINER}" + delete_container_if_present "${POSTGRES_CONTAINER}" + + if resource_exists network "${NETWORK_NAME}"; then + assert_resource_owned network "${NETWORK_NAME}" + info "Deleting network ${NETWORK_NAME}..." + container network delete "${NETWORK_NAME}" >/dev/null + fi + + if [[ "${include_volumes}" == true ]]; then + delete_volume_if_present "${APP_VOLUME}" + delete_volume_if_present "${REDIS_VOLUME}" + delete_volume_if_present "${POSTGRES_VOLUME}" + info "Sub2API stack and persistent data deleted." + else + info "Sub2API stack deleted; persistent volumes were preserved." + fi +} + +main() { + local command=${1-} + if [[ $# -gt 0 ]]; then + shift + fi + + case "${command}" in + init) + [[ $# -eq 0 ]] || { usage; exit 2; } + acquire_lock + cmd_init + ;; + up) + acquire_lock + cmd_up "$@" + ;; + down) + [[ $# -eq 0 ]] || { usage; exit 2; } + acquire_lock + cmd_down + ;; + restart) + [[ $# -eq 0 ]] || { usage; exit 2; } + acquire_lock + cmd_restart + ;; + status) + [[ $# -eq 0 ]] || { usage; exit 2; } + trap cleanup EXIT + cmd_status + ;; + logs) + cmd_logs "$@" + ;; + pull) + [[ $# -eq 0 ]] || { usage; exit 2; } + acquire_lock + cmd_pull + ;; + destroy) + acquire_lock + cmd_destroy "$@" + ;; + help|-h|--help) + usage + ;; + *) + usage + exit 2 + ;; + esac +} + +main "$@" diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index eff4bfb598..4cf759497f 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -20,6 +20,9 @@ server: # Mode: "debug" for development, "release" for production # 运行模式:"debug" 用于开发,"release" 用于生产环境 mode: "release" + # Return Server-Timing for authenticated requests made by the Admin web UI + # 为管理端 Web 页面发出的已认证请求返回 Server-Timing + enable_server_timing: false # Frontend base URL used to generate external links in emails (e.g. password reset) # 用于生成邮件中的外部链接(例如:重置密码链接)的前端基础地址 # Example: "https://example.com" @@ -254,6 +257,10 @@ gateway: # ingress 默认模式:off|ctx_pool|passthrough|http_bridge(仅 mode_router_v2_enabled=true 生效) # 兼容旧值:shared/dedicated 会按 ctx_pool 处理。 ingress_mode_default: ctx_pool + # Close a client WebSocket that stays idle between completed turns (seconds). Set 0 to disable. + ingress_inter_turn_idle_timeout_seconds: 300 + # Limit live client WebSocket ingress sessions per API key across all instances. Set 0 to disable. + max_ingress_connections_per_api_key: 64 # 全局总开关,默认 true;关闭时所有请求保持原有 HTTP/SSE 路由 enabled: true # 按账号类型细分开关 @@ -383,6 +390,9 @@ gateway: # Image stream keepalive interval (seconds), 0=disable; independent from ordinary text streams # 图片流式 keepalive 间隔(秒),0=禁用;独立于普通文本流式 image_stream_keepalive_interval: 10 + # Non-streaming Images JSON keepalive interval (seconds), 0=disable; commits HTTP 200 after the first heartbeat + # 图片非流式 JSON keepalive 间隔(秒),0=禁用;首个心跳后 HTTP 状态会固化为 200 + image_nonstream_keepalive_interval: 0 # Image generation independent concurrency limiter (process-local, default disabled) # 图片生成独立并发限制(进程级,默认关闭;多实例总上限约为实例数×该值) image_concurrency: diff --git a/deploy/docker-compose.dev.yml b/deploy/docker-compose.dev.yml index 6f5b3f56f3..43f5dd3f60 100644 --- a/deploy/docker-compose.dev.yml +++ b/deploy/docker-compose.dev.yml @@ -26,6 +26,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=debug + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} - DATABASE_HOST=postgres - DATABASE_PORT=5432 diff --git a/deploy/docker-compose.local.yml b/deploy/docker-compose.local.yml index 042752e857..5fb161603b 100644 --- a/deploy/docker-compose.local.yml +++ b/deploy/docker-compose.local.yml @@ -51,6 +51,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/deploy/docker-compose.standalone.yml b/deploy/docker-compose.standalone.yml index 2e1d335624..40ed4751d6 100644 --- a/deploy/docker-compose.standalone.yml +++ b/deploy/docker-compose.standalone.yml @@ -37,6 +37,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index 22713c59aa..6aecdcfa5a 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -47,6 +47,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/deploy/tests/apple-container-test.sh b/deploy/tests/apple-container-test.sh new file mode 100755 index 0000000000..a12582104f --- /dev/null +++ b/deploy/tests/apple-container-test.sh @@ -0,0 +1,78 @@ +#!/bin/bash + +set -euo pipefail + +TEST_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +DEPLOY_DIR="$(cd "${TEST_DIR}/.." && pwd)" +SCRIPT="${DEPLOY_DIR}/apple-container.sh" +TEST_ROOT="$(mktemp -d "${TMPDIR:-/tmp}/sub2api-apple-test.XXXXXX")" +STATE_DIR="${TEST_ROOT}/state" +ENV_FILE="${TEST_ROOT}/sub2api.env" + +cleanup() { + rm -rf "${TEST_ROOT}" +} +trap cleanup EXIT + +fail() { + printf 'FAIL: %s\n' "$*" >&2 + exit 1 +} + +assert_exists() { + [[ -e "$1" ]] || fail "Expected path to exist: $1" +} + +assert_missing() { + [[ ! -e "$1" ]] || fail "Expected path to be absent: $1" +} + +export FAKE_CONTAINER_STATE="${STATE_DIR}" +export PATH="${TEST_DIR}/fixtures/bin:${PATH}" +export SUB2API_ENV_FILE="${ENV_FILE}" + +mkdir -p "${STATE_DIR}" + +"${SCRIPT}" init +[[ "$(stat -f '%Lp' "${ENV_FILE}")" == "600" ]] || fail "init did not create a mode-600 env file" +grep -q '^POSTGRES_PASSWORD=change_this_secure_password$' "${ENV_FILE}" && fail "init retained the placeholder password" + +chmod 644 "${ENV_FILE}" +if "${SCRIPT}" up >/dev/null 2>&1; then + fail "up accepted an insecure env file" +fi +chmod 600 "${ENV_FILE}" + +"${SCRIPT}" up +assert_exists "${STATE_DIR}/containers/sub2api-apple" +assert_exists "${STATE_DIR}/containers/sub2api-apple-postgres" +assert_exists "${STATE_DIR}/containers/sub2api-apple-redis" +assert_exists "${STATE_DIR}/running/sub2api-apple" +"${SCRIPT}" status >/dev/null + +"${SCRIPT}" up --recreate +assert_exists "${STATE_DIR}/running/sub2api-apple" +"${SCRIPT}" down +assert_missing "${STATE_DIR}/running/sub2api-apple" +assert_missing "${STATE_DIR}/running/sub2api-apple-postgres" +assert_missing "${STATE_DIR}/running/sub2api-apple-redis" + +"${SCRIPT}" destroy --yes +assert_missing "${STATE_DIR}/containers/sub2api-apple" +assert_missing "${STATE_DIR}/networks/sub2api-apple" +assert_exists "${STATE_DIR}/volumes/sub2api-apple-data" + +"${SCRIPT}" up +"${SCRIPT}" destroy --volumes --yes +assert_missing "${STATE_DIR}/volumes/sub2api-apple-data" +assert_missing "${STATE_DIR}/volumes/sub2api-apple-postgres-data" +assert_missing "${STATE_DIR}/volumes/sub2api-apple-redis-data" + +touch "${STATE_DIR}/system-running" +touch "${STATE_DIR}/containers/sub2api-apple" +touch "${STATE_DIR}/unowned/container/sub2api-apple" +if "${SCRIPT}" status >/dev/null 2>&1; then + fail "status accepted an unowned same-name container" +fi + +printf 'Apple container lifecycle tests passed.\n' diff --git a/deploy/tests/fixtures/bin/container b/deploy/tests/fixtures/bin/container new file mode 100755 index 0000000000..a864111f27 --- /dev/null +++ b/deploy/tests/fixtures/bin/container @@ -0,0 +1,164 @@ +#!/bin/bash + +set -eu + +STATE_DIR="${FAKE_CONTAINER_STATE:?FAKE_CONTAINER_STATE is required}" +mkdir -p \ + "${STATE_DIR}/containers" \ + "${STATE_DIR}/running" \ + "${STATE_DIR}/networks" \ + "${STATE_DIR}/volumes" \ + "${STATE_DIR}/unowned/container" \ + "${STATE_DIR}/unowned/network" \ + "${STATE_DIR}/unowned/volume" + +list_names() { + local directory=$1 + local path + + for path in "${directory}"/*; do + [[ -e "${path}" ]] || continue + basename "${path}" + done +} + +last_argument() { + local value="" + + for value in "$@"; do :; done + printf '%s\n' "${value}" +} + +inspect_resource() { + local resource_type=$1 + local resource_name=$2 + local label_value="apple-container" + local address="192.168.65.4/24" + + if [[ -e "${STATE_DIR}/unowned/${resource_type}/${resource_name}" ]]; then + label_value="other" + fi + case "${resource_name}" in + sub2api-apple-postgres) address="192.168.65.2/24" ;; + sub2api-apple-redis) address="192.168.65.3/24" ;; + esac + + printf '[{"configuration":{"labels":{"org.sub2api.stack":"%s"}},"status":{"networks":[{"ipv4Address":"%s"}]}}]\n' \ + "${label_value}" "${address}" +} + +command=${1-} +if [[ $# -gt 0 ]]; then shift; fi + +case "${command}" in + --version) + echo "container CLI version 1.1.0 (build: release, commit: fake)" + ;; + system) + subcommand=${1-} + case "${subcommand}" in + status) [[ -e "${STATE_DIR}/system-running" ]] ;; + start) touch "${STATE_DIR}/system-running" ;; + stop) rm -f "${STATE_DIR}/system-running" "${STATE_DIR}/running"/* ;; + *) exit 1 ;; + esac + ;; + list) + include_all=false + for argument in "$@"; do + [[ "${argument}" == "--all" || "${argument}" == "-a" ]] && include_all=true + done + if [[ "${include_all}" == true ]]; then + list_names "${STATE_DIR}/containers" + else + list_names "${STATE_DIR}/running" + fi + ;; + network) + subcommand=${1-} + shift || true + case "${subcommand}" in + list) + echo default + list_names "${STATE_DIR}/networks" + ;; + create) touch "${STATE_DIR}/networks/$(last_argument "$@")" ;; + inspect) inspect_resource network "${1}" ;; + delete) rm -f "${STATE_DIR}/networks/${1}" ;; + *) exit 1 ;; + esac + ;; + volume) + subcommand=${1-} + shift || true + case "${subcommand}" in + list) list_names "${STATE_DIR}/volumes" ;; + create) touch "${STATE_DIR}/volumes/$(last_argument "$@")" ;; + inspect) inspect_resource volume "${1}" ;; + delete) rm -f "${STATE_DIR}/volumes/${1}" ;; + *) exit 1 ;; + esac + ;; + image) + subcommand=${1-} + case "${subcommand}" in + inspect|pull) exit 0 ;; + *) exit 1 ;; + esac + ;; + create) + name="" + while [[ $# -gt 0 ]]; do + case "$1" in + --name) + name=$2 + shift 2 + ;; + --label|--network|--platform|--ulimit|--env-file|--volume|--entrypoint|--publish) + shift 2 + ;; + *) + shift + ;; + esac + done + [[ -n "${name}" ]] + touch "${STATE_DIR}/containers/${name}" + ;; + inspect) + inspect_resource container "${1}" + ;; + start) + touch "${STATE_DIR}/running/${1}" + ;; + stop) + for argument in "$@"; do + case "${argument}" in + --time|--signal) skip_next=true ;; + [0-9]*|SIG*) ;; + *) rm -f "${STATE_DIR}/running/${argument}" ;; + esac + done + ;; + delete) + for argument in "$@"; do + case "${argument}" in + --force|-f) ;; + *) + rm -f "${STATE_DIR}/running/${argument}" + rm -f "${STATE_DIR}/containers/${argument}" + ;; + esac + done + ;; + exec) + echo 1 + ;; + logs|copy) + exit 0 + ;; + *) + echo "Unsupported fake container command: ${command} $*" >&2 + exit 1 + ;; +esac diff --git a/deploy/tests/fixtures/bin/curl b/deploy/tests/fixtures/bin/curl new file mode 100755 index 0000000000..5e611b474f --- /dev/null +++ b/deploy/tests/fixtures/bin/curl @@ -0,0 +1,4 @@ +#!/bin/bash + +set -eu +printf '{"status":"ok"}\n' diff --git a/frontend/src/api/__tests__/admin.grok.spec.ts b/frontend/src/api/__tests__/admin.grok.spec.ts new file mode 100644 index 0000000000..7560ce443f --- /dev/null +++ b/frontend/src/api/__tests__/admin.grok.spec.ts @@ -0,0 +1,37 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { post } = vi.hoisted(() => ({ + post: vi.fn(), +})) + +vi.mock('@/api/client', () => ({ + apiClient: { post }, +})) + +import { createFromSSO, getGrokSSOImportTimeout } from '@/api/admin/grok' + +describe('admin Grok SSO import API', () => { + beforeEach(() => { + post.mockReset() + post.mockResolvedValue({ data: { created: [], failed: [] } }) + }) + + it.each([ + [1, 180_000], + [3, 180_000], + [4, 270_000], + [7, 360_000], + ])('uses a timeout sized for %i keys', async (keyCount, expectedTimeout) => { + expect(getGrokSSOImportTimeout(keyCount)).toBe(expectedTimeout) + + await createFromSSO({ + sso_tokens: Array.from({ length: keyCount }, (_, index) => `sso-${index + 1}`), + }) + + expect(post).toHaveBeenCalledWith( + '/admin/grok/sso-to-oauth', + expect.objectContaining({ sso_tokens: expect.any(Array) }), + { timeout: expectedTimeout }, + ) + }) +}) diff --git a/frontend/src/api/__tests__/adminUIRequest.spec.ts b/frontend/src/api/__tests__/adminUIRequest.spec.ts new file mode 100644 index 0000000000..9064a52f1a --- /dev/null +++ b/frontend/src/api/__tests__/adminUIRequest.spec.ts @@ -0,0 +1,38 @@ +import { describe, expect, it } from 'vitest' + +import { + ADMIN_UI_REQUEST_HEADER, + shouldMarkAdminUIRequest, +} from '@/api/adminUIRequest' + +describe('Admin UI request marker', () => { + it('uses the stable request header name', () => { + expect(ADMIN_UI_REQUEST_HEADER).toBe('X-Admin-UI-Request') + }) + + it.each([ + '/admin', + '/admin/users', + '/api/v1/admin', + '/api/v1/admin/accounts?status=active', + 'https://api.example.test/api/v1/admin/dashboard', + ])('marks Admin API request %s before page navigation', (requestURL) => { + expect(shouldMarkAdminUIRequest(requestURL, '/login')).toBe(true) + }) + + it.each(['/keys', '/groups/available', '/auth/me', '/announcements'])( + 'marks shared request %s while an Admin page is active', + (requestURL) => { + expect(shouldMarkAdminUIRequest(requestURL, '/admin/dashboard')).toBe(true) + } + ) + + it.each([ + ['/keys', '/dashboard'], + ['/api/v1/administer', '/dashboard'], + ['/keys', '/administrator'], + ['', '/'], + ])('does not mark request %s on page %s', (requestURL, pagePath) => { + expect(shouldMarkAdminUIRequest(requestURL, pagePath)).toBe(false) + }) +}) diff --git a/frontend/src/api/__tests__/client.spec.ts b/frontend/src/api/__tests__/client.spec.ts index a0a05410d4..b275cca34b 100644 --- a/frontend/src/api/__tests__/client.spec.ts +++ b/frontend/src/api/__tests__/client.spec.ts @@ -12,6 +12,7 @@ describe('API Client', () => { beforeEach(async () => { localStorage.clear() + window.history.replaceState({}, '', '/') // 每次测试重新导入以获取干净的模块状态 vi.resetModules() const mod = await import('@/api/client') @@ -120,6 +121,55 @@ describe('API Client', () => { const config = adapter.mock.calls[0][0] expect(config.withCredentials).toBe(true) }) + + it('Admin API 在进入管理页面前也带 Admin UI 标记', async () => { + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/admin/users') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBe('1') + }) + + it('管理页面调用共享 API 时带 Admin UI 标记', async () => { + window.history.replaceState({}, '', '/admin/dashboard') + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/groups/available') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBe('1') + }) + + it('普通用户页面调用共享 API 时不带 Admin UI 标记', async () => { + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/groups/available') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBeFalsy() + }) }) // --- 响应拦截器 --- diff --git a/frontend/src/api/admin/channelMonitor.ts b/frontend/src/api/admin/channelMonitor.ts index 0b9c62231c..de605351e3 100644 --- a/frontend/src/api/admin/channelMonitor.ts +++ b/frontend/src/api/admin/channelMonitor.ts @@ -5,7 +5,7 @@ import { apiClient } from '../client' -export type Provider = 'openai' | 'anthropic' | 'gemini' +export type Provider = 'openai' | 'anthropic' | 'gemini' | 'grok' export type MonitorStatus = 'operational' | 'degraded' | 'failed' | 'error' export type BodyOverrideMode = 'off' | 'merge' | 'replace' export type APIMode = 'chat_completions' | 'responses' diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts index c0055d4dcc..f59c05c1d7 100644 --- a/frontend/src/api/admin/grok.ts +++ b/frontend/src/api/admin/grok.ts @@ -4,6 +4,9 @@ */ import { apiClient } from '../client' +import type { GrokBillingSummary, GrokQuotaWindow, WindowStats } from '@/types' + +export type { GrokBillingSummary, GrokQuotaWindow } from '@/types' export interface GrokAuthUrlResponse { auth_url: string @@ -34,16 +37,49 @@ export interface GrokTokenInfo { scope?: string client_id?: string email?: string + sub?: string + team_id?: string subscription_tier?: string entitlement_status?: string [key: string]: unknown } -export interface GrokQuotaWindow { - limit?: number | null - remaining?: number | null - reset_unix?: number | null - reset_at?: string | null +export interface GrokSSOToOAuthRequest { + sso_tokens: string[] + name?: string + notes?: string | null + proxy_id?: number | null + group_ids?: number[] + credentials?: Record + extra?: Record + concurrency?: number + load_factor?: number + priority?: number + rate_multiplier?: number + expires_at?: number | null + auto_pause_on_expired?: boolean +} + +export interface GrokSSOToOAuthItemResult { + index: number + name?: string + email?: string + account?: unknown + error?: string +} + +export interface GrokSSOToOAuthResponse { + created: GrokSSOToOAuthItemResult[] + failed: GrokSSOToOAuthItemResult[] +} + +const GROK_SSO_IMPORT_CONCURRENCY = 3 +const GROK_SSO_IMPORT_TIMEOUT_PER_BATCH_MS = 90_000 +const GROK_SSO_IMPORT_TIMEOUT_BUFFER_MS = 90_000 + +export function getGrokSSOImportTimeout(keyCount: number): number { + const batches = Math.ceil(Math.max(1, keyCount) / GROK_SSO_IMPORT_CONCURRENCY) + return batches * GROK_SSO_IMPORT_TIMEOUT_PER_BATCH_MS + GROK_SSO_IMPORT_TIMEOUT_BUFFER_MS } export interface GrokQuotaSnapshot { @@ -62,13 +98,19 @@ export interface GrokQuotaSnapshot { } export interface GrokQuotaProbeResult { - source: 'active_probe' - model: string + source: 'active_probe' | 'billing_probe' | 'hybrid_probe' + model?: string + billing?: GrokBillingSummary | null snapshot?: GrokQuotaSnapshot | null + local_usage_24h?: WindowStats | null + local_usage_7d?: WindowStats | null + local_usage_monthly?: WindowStats | null status_code?: number headers_observed: boolean reset_supported: boolean fetched_at: number + persisted?: boolean + probe_error?: string } export interface GrokQuotaResetResult { @@ -119,4 +161,13 @@ export async function resetQuota(id: number): Promise { return data } -export default { generateAuthUrl, exchangeCode, refreshGrokToken, queryQuota, resetQuota } +export async function createFromSSO(payload: GrokSSOToOAuthRequest): Promise { + const { data } = await apiClient.post( + '/admin/grok/sso-to-oauth', + payload, + { timeout: getGrokSSOImportTimeout(payload.sso_tokens.length) } + ) + return data +} + +export default { generateAuthUrl, exchangeCode, refreshGrokToken, queryQuota, resetQuota, createFromSSO } diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts index c7cbc64a4b..3284ef7e15 100644 --- a/frontend/src/api/admin/ops.ts +++ b/frontend/src/api/admin/ops.ts @@ -828,6 +828,7 @@ export interface OpsRuntimeLogConfig { export interface OpsSystemLog { id: number created_at: string + host: string level: string component: string message: string @@ -849,6 +850,7 @@ export interface OpsSystemLogQuery { time_range?: '5m' | '30m' | '1h' | '6h' | '24h' | '7d' | '30d' start_time?: string end_time?: string + host?: string level?: string component?: string request_id?: string @@ -864,6 +866,7 @@ export interface OpsSystemLogQuery { export interface OpsSystemLogCleanupRequest { start_time?: string end_time?: string + host?: string level?: string component?: string request_id?: string diff --git a/frontend/src/api/adminUIRequest.ts b/frontend/src/api/adminUIRequest.ts new file mode 100644 index 0000000000..2d60e2987d --- /dev/null +++ b/frontend/src/api/adminUIRequest.ts @@ -0,0 +1,27 @@ +export const ADMIN_UI_REQUEST_HEADER = 'X-Admin-UI-Request' + +function isAdminPath(path: string): boolean { + return ( + path === '/admin' || + path.startsWith('/admin/') || + path === '/api/v1/admin' || + path.startsWith('/api/v1/admin/') + ) +} + +function requestPath(rawURL: string): string { + const value = rawURL.trim() + if (!value) return '' + try { + const origin = typeof window !== 'undefined' ? window.location.origin : 'http://localhost' + return new URL(value, origin).pathname + } catch { + return value.split(/[?#]/, 1)[0] + } +} + +export function shouldMarkAdminUIRequest(requestURL: string, pagePath?: string): boolean { + const currentPath = + pagePath ?? (typeof window !== 'undefined' ? window.location.pathname : '') + return isAdminPath(requestPath(requestURL)) || isAdminPath(currentPath) +} diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 5df969f188..a2b4d2f650 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -6,6 +6,7 @@ import axios, { AxiosInstance, AxiosError, InternalAxiosRequestConfig, AxiosResponse } from 'axios' import type { ApiResponse } from '@/types' import { getLocale } from '@/i18n' +import { ADMIN_UI_REQUEST_HEADER, shouldMarkAdminUIRequest } from './adminUIRequest' import { getAPIBaseURL } from './url' export { buildApiUrl, buildGatewayUrl } from './url' @@ -74,6 +75,10 @@ apiClient.interceptors.request.use( config.params.timezone = getUserTimezone() } + if (config.headers && shouldMarkAdminUIRequest(String(config.url || ''))) { + config.headers[ADMIN_UI_REQUEST_HEADER] = '1' + } + return config }, (error) => { diff --git a/frontend/src/api/payment.ts b/frontend/src/api/payment.ts index ab83c55d94..e18508c535 100644 --- a/frontend/src/api/payment.ts +++ b/frontend/src/api/payment.ts @@ -7,7 +7,6 @@ import { apiClient } from './client' import type { PaymentConfig, SubscriptionPlan, - PaymentChannel, MethodLimitsResponse, CheckoutInfoResponse, CreateOrderRequest, @@ -35,11 +34,6 @@ export const paymentAPI = { return apiClient.get('/payment/plans') }, - /** Get available payment channels */ - getChannels() { - return apiClient.get('/payment/channels') - }, - /** Get all checkout page data in a single call */ getCheckoutInfo() { return apiClient.get('/payment/checkout-info') diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 81d97efb9c..301b512c25 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -382,17 +382,35 @@ + +
@@ -407,7 +425,7 @@
{{ grokQuotaStatusLine }}
- +
-
@@ -600,6 +618,7 @@ import { ref, computed, onMounted, onBeforeUnmount, onUnmounted, watch } from 'vue' import { useI18n } from 'vue-i18n' import { adminAPI } from '@/api/admin' +import type { GrokQuotaProbeResult } from '@/api/admin/grok' import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from '@/types' import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { enqueueUsageRequest } from '@/utils/usageLoadQueue' @@ -612,6 +631,8 @@ import GrokQuotaProbeCell from './GrokQuotaProbeCell.vue' // Module-level cache shared across all AccountUsageCell instances const _usageCache = new Map() const USAGE_CACHE_TTL = 5 * 60 * 1000 // 5 minutes +// xAI Free billing exposes a window without usage_percent, so estimate it from local tokens. +const GROK_FREE_TOKEN_LIMIT = 2_000_000 const props = withDefaults( defineProps<{ @@ -1036,18 +1057,67 @@ interface GrokQuotaBarInfo { const makeGrokQuotaBar = (quota?: { limit?: number | null; remaining?: number | null; reset_at?: string | null } | null): GrokQuotaBarInfo | null => { if (!quota || quota.limit == null || quota.remaining == null || quota.limit <= 0) return null - const used = Math.max(0, quota.limit - quota.remaining) + const remaining = Math.min(quota.limit, Math.max(0, quota.remaining)) return { - utilization: (used / quota.limit) * 100, + utilization: (remaining / quota.limit) * 100, resetsAt: quota.reset_at || null } } const grokRequestQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_request_quota)) const grokTokenQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_token_quota)) +const grokBilling = computed(() => usageInfo.value?.grok_billing || null) +const grokWeeklyBillingBar = computed((): GrokQuotaBarInfo | null => { + const billing = grokBilling.value + if (billing?.period_type?.toLowerCase() !== 'weekly' || billing.usage_percent == null) { + return null + } + return { + utilization: Math.min(100, Math.max(0, billing.usage_percent)), + resetsAt: billing.period_end || null + } +}) +const grokPlanLabelIsFree = (value: string) => value.includes('free') || value.includes('basic') +const grokPlanLabelIsPaid = (value: string) => { + return value !== '' && !grokPlanLabelIsFree(value) && !value.includes('unknown') +} +const grokIsFree = computed(() => { + if (props.account.platform !== 'grok' || props.account.type !== 'oauth') return false + const billing = grokBilling.value + if ( + billing?.usage_percent != null || + billing?.used_percent != null || + (billing?.monthly_limit_cents != null && billing.monthly_limit_cents > 0) + ) return false + + const plan = (billing?.plan || '').trim().toLowerCase() + const tier = (usageInfo.value?.subscription_tier || '').trim().toLowerCase() + const entitlement = (usageInfo.value?.grok_entitlement_status || '').toLowerCase() + if (grokPlanLabelIsPaid(plan) || grokPlanLabelIsPaid(tier)) return false + if ( + grokPlanLabelIsFree(plan) || + grokPlanLabelIsFree(tier) || + grokPlanLabelIsFree(entitlement) + ) return true + return billing != null +}) +const grokFreeQuotaUsage = computed(() => usageInfo.value?.grok_local_usage_24h || null) +const grokLocalUsage = computed(() => { + if (grokIsFree.value) return grokFreeQuotaUsage.value + return props.todayStats || + usageInfo.value?.grok_local_usage || + usageInfo.value?.grok_local_usage_7d || + usageInfo.value?.grok_local_usage_monthly || + null +}) +const grokFreeTokenBar = computed(() => { + if (!grokIsFree.value || !grokFreeQuotaUsage.value) return null + const used = Math.max(0, grokFreeQuotaUsage.value.tokens || 0) + return { utilization: Math.min(100, (used / GROK_FREE_TOKEN_LIMIT) * 100) } +}) const grokQuotaUnknown = computed(() => { if (props.account.platform !== 'grok') return false - if (grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false + if (grokBilling.value || grokFreeTokenBar.value || grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false return usageInfo.value?.grok_quota_snapshot_state !== 'observed' }) const grokQuotaUnknownLabel = computed(() => { @@ -1078,7 +1148,6 @@ const grokQuotaStatusLine = computed(() => { } return parts.length > 0 ? parts.join(' | ') : null }) -const grokLocalUsage = computed(() => usageInfo.value?.grok_local_usage || props.todayStats || null) const grokEntitlementLabel = computed(() => { const status = (usageInfo.value?.grok_entitlement_status || '').trim() return status || null @@ -1281,6 +1350,35 @@ const loadActiveUsage = async () => { } } +const handleGrokProbed = (result: GrokQuotaProbeResult) => { + const current = usageInfo.value + if (!current) return + const snapshot = result.snapshot + const merged: AccountUsageInfo = { + ...current, + grok_billing: result.billing ?? current.grok_billing, + grok_local_usage_24h: result.local_usage_24h ?? current.grok_local_usage_24h, + grok_local_usage_7d: result.local_usage_7d ?? current.grok_local_usage_7d, + grok_local_usage_monthly: result.local_usage_monthly ?? current.grok_local_usage_monthly, + grok_request_quota: snapshot?.requests ?? current.grok_request_quota, + grok_token_quota: snapshot?.tokens ?? current.grok_token_quota, + grok_retry_after_seconds: snapshot?.retry_after_seconds ?? current.grok_retry_after_seconds, + grok_entitlement_status: snapshot?.entitlement_status || current.grok_entitlement_status, + grok_quota_snapshot_state: result.billing + ? 'billing_observed' + : snapshot?.headers_observed + ? 'observed' + : current.grok_quota_snapshot_state, + grok_last_quota_probe_at: result.billing?.fetched_at ?? snapshot?.last_probe_at ?? current.grok_last_quota_probe_at, + grok_last_headers_seen_at: snapshot?.last_headers_seen_at ?? current.grok_last_headers_seen_at, + grok_last_status_code: result.status_code ?? snapshot?.status_code ?? current.grok_last_status_code, + error: result.billing || snapshot ? undefined : current.error, + error_code: result.billing || snapshot ? undefined : current.error_code + } + usageInfo.value = merged + _usageCache.set(props.account.id, { data: merged, ts: Date.now() }) +} + // ===== API Key quota progress bars ===== interface QuotaBarInfo { diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 67750e4e34..ea9deb7ba6 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -50,7 +50,7 @@ - +
@@ -381,10 +381,34 @@ {{ t('admin.accounts.types.grokOauth') }}
+ +
-

- {{ t('admin.accounts.oauth.grok.oauthOnlyHint') }} -

@@ -1087,10 +1111,12 @@ ? 'https://api.openai.com' : form.platform === 'gemini' ? 'https://generativelanguage.googleapis.com' - : 'https://api.anthropic.com' + : form.platform === 'grok' + ? 'https://api.x.ai/v1' + : 'https://api.anthropic.com' " /> -

{{ baseUrlHint }}

+

{{ baseUrlHint }}

@@ -1104,10 +1130,12 @@ ? 'sk-proj-...' : form.platform === 'gemini' ? 'AIza...' - : 'sk-ant-...' + : form.platform === 'grok' + ? 'xai-...' + : 'sk-ant-...' " /> -

{{ apiKeyHint }}

+

{{ apiKeyHint }}

@@ -1938,7 +1966,7 @@
@@ -2796,6 +2824,38 @@
+
+
+
+ +

+ {{ t('admin.accounts.openai.longContextBillingDesc') }} +

+
+ +
+
+
@@ -3464,6 +3528,7 @@ interface OAuthFlowExposed { sessionToken: string codexSession: string codexPAT: string + ssoCookie: string inputMethod: AuthInputMethod reset: () => void } @@ -3483,14 +3548,14 @@ const oauthStepTitle = computed(() => { const baseUrlHint = computed(() => { if (form.platform === 'openai') return t('admin.accounts.openai.baseUrlHint') if (form.platform === 'gemini') return t('admin.accounts.gemini.baseUrlHint') - if (form.platform === 'grok') return t('admin.accounts.grok.baseUrlHint') + if (form.platform === 'grok') return '' return t('admin.accounts.baseUrlHint') }) const apiKeyHint = computed(() => { if (form.platform === 'openai') return t('admin.accounts.openai.apiKeyHint') if (form.platform === 'gemini') return t('admin.accounts.gemini.apiKeyHint') - if (form.platform === 'grok') return t('admin.accounts.grok.apiKeyHint') + if (form.platform === 'grok') return '' return t('admin.accounts.apiKeyHint') }) @@ -3648,6 +3713,8 @@ const fillHeaderOverrideTemplate = () => { const interceptWarmupRequests = ref(false) const autoPauseOnExpired = ref(true) const openaiPassthroughEnabled = ref(false) +const openAILongContextBillingEnabled = ref(false) +const openAILongContextBillingTouched = ref(false) const openAICompactMode = ref('auto') const openAIResponsesMode = ref('auto') const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings']) @@ -3660,6 +3727,11 @@ const anthropicPassthroughEnabled = ref(false) const anthropicAPIKeyAuthScheme = ref('x_api_key') const webSearchEmulationMode = ref('default') const webSearchGlobalEnabled = ref(false) + +const toggleOpenAILongContextBilling = () => { + openAILongContextBillingEnabled.value = !openAILongContextBillingEnabled.value + openAILongContextBillingTouched.value = true +} const { globalEnabled: quotaNotifyGlobalEnabled, state: quotaNotifyState, @@ -3948,6 +4020,8 @@ const isOAuthFlow = computed(() => { return accountCategory.value === 'oauth-based' }) +const isGrokSSOInputMethod = computed(() => form.platform === 'grok' && oauthFlowRef.value?.inputMethod === 'sso_cookie') + const isManualInputMethod = computed(() => { return oauthFlowRef.value?.inputMethod === 'manual' }) @@ -4502,6 +4576,8 @@ const resetForm = () => { interceptWarmupRequests.value = false autoPauseOnExpired.value = true openaiPassthroughEnabled.value = false + openAILongContextBillingEnabled.value = false + openAILongContextBillingTouched.value = false openAICompactMode.value = 'auto' openAIResponsesMode.value = 'auto' openAIEndpointCapabilities.value = ['chat_completions', 'embeddings'] @@ -4584,6 +4660,7 @@ const buildOpenAIExtra = (base?: Record): Record): Record 0 ? extra : undefined } +const buildOpenAICodexImportExtra = (): Record | undefined => { + const extra = buildOpenAIExtra() + if (!extra) { + return undefined + } + if (!openAILongContextBillingTouched.value) { + delete extra.openai_long_context_billing_enabled + } + return Object.keys(extra).length > 0 ? extra : undefined +} + const buildAnthropicExtra = (base?: Record): Record | undefined => { if (form.platform !== 'anthropic' || accountCategory.value !== 'apikey') { return base @@ -4738,7 +4826,7 @@ const handleVertexServiceAccountDrop = async (event: DragEvent) => { const handleSubmit = async () => { // For OAuth-based type, handle OAuth flow (goes to step 2) if (isOAuthFlow.value) { - if (!form.name.trim()) { + if (!isGrokSSOInputMethod.value && !form.name.trim()) { appStore.showError(t('admin.accounts.pleaseEnterAccountName')) return } @@ -4887,7 +4975,9 @@ const handleSubmit = async () => { ? 'https://api.openai.com' : form.platform === 'gemini' ? 'https://generativelanguage.googleapis.com' - : 'https://api.anthropic.com' + : form.platform === 'grok' + ? 'https://api.x.ai/v1' + : 'https://api.anthropic.com' // Build credentials with optional model mapping const credentials: Record = { @@ -5174,6 +5264,76 @@ const handleGrokValidateRT = async (refreshTokenInput: string) => { } } +const handleGrokImportSSO = async (ssoInput: string) => { + // Align with OpenAI/Grok RT batch import: one token per line, no client-side dedupe. + const ssoTokens = ssoInput + .split('\n') + .map((token) => token.trim()) + .filter((token) => token) + if (ssoTokens.length === 0) return + + grokOAuth.loading.value = true + grokOAuth.error.value = '' + + const credentials: Record = {} + const modelMapping = buildModelMappingObject(modelRestrictionMode.value, allowedModels.value, modelMappings.value) + if (modelMapping) { + credentials.model_mapping = modelMapping + } + if (!applyTempUnschedConfig(credentials)) { + grokOAuth.loading.value = false + return + } + + try { + const result = await adminAPI.grok.createFromSSO({ + sso_tokens: ssoTokens, + name: form.name || undefined, + notes: form.notes || undefined, + proxy_id: form.proxy_id, + group_ids: form.group_ids, + credentials, + concurrency: form.concurrency, + load_factor: form.load_factor ?? undefined, + priority: form.priority, + rate_multiplier: form.rate_multiplier, + expires_at: form.expires_at, + auto_pause_on_expired: autoPauseOnExpired.value + }) + + const successCount = result.created?.length || 0 + const failedCount = result.failed?.length || 0 + if (successCount > 0 && failedCount === 0) { + appStore.showSuccess( + ssoTokens.length > 1 + ? t('admin.accounts.oauth.batchSuccess', { count: successCount }) + : t('admin.accounts.accountCreated') + ) + emit('created') + handleClose() + } else if (successCount > 0 && failedCount > 0) { + // Same as OpenAI/Grok RT: keep input, show failures, refresh list. + appStore.showWarning( + t('admin.accounts.oauth.batchPartialSuccess', { success: successCount, failed: failedCount }) + ) + grokOAuth.error.value = (result.failed || []) + .map((item) => `#${item.index}: ${item.error || 'Unknown error'}`) + .join('\n') + emit('created') + } else { + grokOAuth.error.value = (result.failed || []) + .map((item) => `#${item.index}: ${item.error || 'Unknown error'}`) + .join('\n') || t('admin.accounts.oauth.grok.failedToConvertSSO') + appStore.showError(t('admin.accounts.oauth.batchFailed')) + } + } catch (error: any) { + grokOAuth.error.value = error.response?.data?.detail || error.message || t('admin.accounts.oauth.grok.failedToConvertSSO') + appStore.showError(grokOAuth.error.value) + } finally { + grokOAuth.loading.value = false + } +} + // OpenAI OAuth 授权码兑换 const handleOpenAIExchange = async (authCode: string) => { const oauthClient = openaiOAuth @@ -5302,7 +5462,7 @@ const handleOpenAIImportCodexSession = async (content: string) => { oauthClient.error.value = '' try { - const extra = buildOpenAIExtra() + const extra = buildOpenAICodexImportExtra() const result = await adminAPI.accounts.importCodexSession({ content: trimmed, name: form.name, @@ -5380,7 +5540,7 @@ const handleOpenAIImportCodexPAT = async (accessToken: string) => { oauthClient.error.value = '' try { - const extra = buildOpenAIExtra() + const extra = buildOpenAICodexImportExtra() await adminAPI.accounts.createOpenAICodexPAT({ access_token: trimmed, name: form.name, diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 9b5ab82fc5..cfc2fed151 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -41,10 +41,12 @@ ? 'https://generativelanguage.googleapis.com' : account.platform === 'antigravity' ? 'https://cloudcode-pa.googleapis.com' - : 'https://api.anthropic.com' + : account.platform === 'grok' + ? 'https://api.x.ai/v1' + : 'https://api.anthropic.com' " /> -

{{ baseUrlHint }}

+

{{ baseUrlHint }}

@@ -63,7 +65,9 @@ ? 'AIza...' : account.platform === 'antigravity' ? 'sk-...' - : 'sk-ant-...' + : account.platform === 'grok' + ? 'xai-...' + : 'sk-ant-...' " />

{{ t('admin.accounts.leaveEmptyToKeep') }}

@@ -1782,7 +1786,39 @@ />
- + +
+
+
+ +

+ {{ t('admin.accounts.openai.longContextBillingDesc') }} +

+
+ +
+
+
+ +
+
+
+ +

+ {{ t('admin.accounts.openai.planTypeDesc') }} +

+
+
+ +
+ +
+
+

+ {{ t(getOAuthKey('ssoCookieDesc')) }} +

+ +
+ + +

+ {{ t(getOAuthKey('ssoCookieHint')) }} +

+
+ +
+

+ {{ error }} +

+
+ + +
+
+
(), { showAccessTokenOption: false, showCodexSessionImportOption: false, showCodexPatOption: false, + showSsoOption: false, + showManualOption: true, + initialInputMethod: 'manual', platform: 'anthropic', showProjectId: true }) @@ -771,6 +863,7 @@ const emit = defineEmits<{ 'import-access-token': [accessToken: string] 'import-codex-session': [content: string] 'import-codex-pat': [accessToken: string] + 'import-sso': [content: string] 'update:inputMethod': [method: AuthInputMethod] }>() @@ -807,19 +900,31 @@ const oauthImportantNotice = computed(() => { }) // Local state -const inputMethod = ref(props.showCookieOption ? 'manual' : 'manual') +const inputMethod = ref(props.initialInputMethod) const authCodeInput = ref('') const sessionKeyInput = ref('') const refreshTokenInput = ref('') const sessionTokenInput = ref('') const codexSessionInput = ref('') const codexPATInput = ref('') +const ssoCookieInput = ref('') const showHelpDialog = ref(false) const oauthState = ref('') const projectId = ref('') -// Computed: show method selection when either cookie or refresh token option is enabled -const showMethodSelection = computed(() => props.showCookieOption || props.showRefreshTokenOption || props.showMobileRefreshTokenOption || props.showSessionTokenOption || props.showAccessTokenOption || props.showCodexSessionImportOption || props.showCodexPatOption) +// Computed: show method selection only when there is something to choose. +const methodOptionCount = computed(() => [ + props.showManualOption, + props.showCookieOption, + props.showRefreshTokenOption, + props.showMobileRefreshTokenOption, + props.showSessionTokenOption, + props.showAccessTokenOption, + props.showCodexSessionImportOption, + props.showCodexPatOption, + props.showSsoOption +].filter(Boolean).length) +const showMethodSelection = computed(() => methodOptionCount.value > 1) // Clipboard const { copied, copyToClipboard } = useClipboard() @@ -850,7 +955,18 @@ const parsedCodexSessionCount = computed(() => { .filter((item) => item).length }) +const parsedSSOCount = computed(() => { + return ssoCookieInput.value + .split('\n') + .map((item) => item.trim()) + .filter((item) => item).length +}) + // Watchers +watch(() => props.initialInputMethod, (newVal) => { + inputMethod.value = newVal +}) + watch(inputMethod, (newVal) => { emit('update:inputMethod', newVal) }) @@ -933,6 +1049,12 @@ const handleImportCodexPAT = () => { } } +const handleImportSSO = () => { + if (ssoCookieInput.value.trim()) { + emit('import-sso', ssoCookieInput.value.trim()) + } +} + // Expose methods and state defineExpose({ authCode: authCodeInput, @@ -943,6 +1065,7 @@ defineExpose({ sessionToken: sessionTokenInput, codexSession: codexSessionInput, codexPAT: codexPATInput, + ssoCookie: ssoCookieInput, inputMethod, reset: () => { authCodeInput.value = '' @@ -953,7 +1076,8 @@ defineExpose({ sessionTokenInput.value = '' codexSessionInput.value = '' codexPATInput.value = '' - inputMethod.value = 'manual' + ssoCookieInput.value = '' + inputMethod.value = props.initialInputMethod showHelpDialog.value = false } }) diff --git a/frontend/src/components/account/UsageProgressBar.vue b/frontend/src/components/account/UsageProgressBar.vue index 6a69357318..2f8b9a88a7 100644 --- a/frontend/src/components/account/UsageProgressBar.vue +++ b/frontend/src/components/account/UsageProgressBar.vue @@ -69,6 +69,7 @@ const props = defineProps<{ color: 'indigo' | 'emerald' | 'purple' | 'amber' windowStats?: WindowStats | null showNowWhenIdle?: boolean + remainingCapacity?: boolean }>() const { t } = useI18n() @@ -109,6 +110,14 @@ const labelClass = computed(() => { // Progress bar color based on utilization const barClass = computed(() => { + if (props.remainingCapacity) { + if (props.utilization <= 20) { + return 'bg-red-500' + } else if (props.utilization <= 50) { + return 'bg-amber-500' + } + return 'bg-green-500' + } if (props.utilization >= 100) { return 'bg-red-500' } else if (props.utilization >= 80) { @@ -120,6 +129,14 @@ const barClass = computed(() => { // Text color based on utilization const textClass = computed(() => { + if (props.remainingCapacity) { + if (props.utilization <= 20) { + return 'text-red-600 dark:text-red-400' + } else if (props.utilization <= 50) { + return 'text-amber-600 dark:text-amber-400' + } + return 'text-gray-600 dark:text-gray-400' + } if (props.utilization >= 100) { return 'text-red-600 dark:text-red-400' } else if (props.utilization >= 80) { @@ -131,12 +148,16 @@ const textClass = computed(() => { // Bar width (capped at 100%) const barWidth = computed(() => { - return `${Math.min(props.utilization, 100)}%` + return `${Math.min(Math.max(props.utilization, 0), 100)}%` }) // Display percentage (cap at 999% for readability) const displayPercent = computed(() => { - const percent = Math.round(props.utilization) + const percent = Math.round( + props.remainingCapacity + ? Math.min(Math.max(props.utilization, 0), 100) + : props.utilization + ) return percent > 999 ? '>999%' : `${percent}%` }) diff --git a/frontend/src/components/account/__tests__/AccountStatusIndicator.spec.ts b/frontend/src/components/account/__tests__/AccountStatusIndicator.spec.ts index f758e6b0f6..545c2abac4 100644 --- a/frontend/src/components/account/__tests__/AccountStatusIndicator.spec.ts +++ b/frontend/src/components/account/__tests__/AccountStatusIndicator.spec.ts @@ -13,6 +13,14 @@ vi.mock('vue-i18n', async () => { } }) +vi.mock('@/utils/format', async () => { + const actual = await vi.importActual('@/utils/format') + return { + ...actual, + formatCountdown: () => '1h' + } +}) + function makeAccount(overrides: Partial): Account { return { id: 1, @@ -43,6 +51,31 @@ function makeAccount(overrides: Partial): Account { } describe('AccountStatusIndicator', () => { + it('Grok 账号额度限流时显示自动恢复时间而非临时不可调度', () => { + const wrapper = mount(AccountStatusIndicator, { + props: { + account: makeAccount({ + id: 5, + name: 'grok-free-1', + platform: 'grok', + rate_limited_at: '2026-07-11T12:00:00Z', + rate_limit_reset_at: '2099-07-11T13:00:00Z', + temp_unschedulable_until: '2099-07-11T12:30:00Z', + temp_unschedulable_reason: 'legacy grok rate limited' + }) + }, + global: { + stubs: { + Icon: true + } + } + }) + + expect(wrapper.find('.badge-warning').text()).toBe('admin.accounts.status.rateLimited') + expect(wrapper.text()).toContain('admin.accounts.status.rateLimitedAutoResume') + expect(wrapper.text()).not.toContain('admin.accounts.status.tempUnschedulable') + }) + it('模型限流 + overages 启用 + 无 AICredits key → 显示 ⚡ (credits_active)', () => { const wrapper = mount(AccountStatusIndicator, { props: { diff --git a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts index 2df3cc07e2..7c2fe48a56 100644 --- a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts +++ b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts @@ -566,7 +566,7 @@ describe('AccountUsageCell', () => { expect(badges.some(node => node.attributes('title') === 'usage.userBilled')).toBe(true) }) - it('Grok OAuth 会展示本地 user billed 用量并保留超限百分比', async () => { + it('Grok OAuth 会展示本地 user billed 用量并把耗尽配额显示为 0% 剩余', async () => { getUsage.mockResolvedValue({ grok_local_usage: { requests: 4, @@ -611,13 +611,440 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).toContain('1.2K') expect(wrapper.text()).toContain('A $0.12') expect(wrapper.text()).toContain('U $0.34') - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokRequests|120|2026-07-09T16:00:00Z') + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokRequests|0|2026-07-09T16:00:00Z') const badges = wrapper.findAll('span[title]') expect(badges.some(node => node.attributes('title') === 'usage.accountBilled')).toBe(true) expect(badges.some(node => node.attributes('title') === 'usage.userBilled')).toBe(true) }) + it('Grok OAuth 配额条按剩余容量显示 100% 满格和 25% 低量', async () => { + getUsage.mockResolvedValue({ + grok_request_quota: { + limit: 100, + remaining: 100, + reset_at: '2026-07-09T16:00:00Z' + }, + grok_token_quota: { + limit: 1000, + remaining: 250, + reset_at: '2026-07-09T16:00:00Z' + }, + grok_quota_snapshot_state: 'observed' + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ + id: 4073, + platform: 'grok', + type: 'oauth', + extra: {} + }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt', 'color', 'remainingCapacity'], + template: '
{{ label }}|{{ utilization }}|{{ remainingCapacity }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokRequests|100|true') + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25|true') + }) + + it('Grok OAuth uses the official weekly billing percentage when available', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: 37, + period_end: '2026-07-16T03:25:00Z', + plan: 'SuperGrok' + }, + grok_local_usage: { + requests: 5, + tokens: 2_200_000, + cost: 4.42, + standard_cost: 4.42, + user_cost: 0.44 + }, + grok_request_quota: { limit: 100, remaining: 100 }, + grok_token_quota: { limit: 2_000_000, remaining: 2_000_000 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4201, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt', 'remainingCapacity'], + template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}|{{ remainingCapacity }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('7d|37|2026-07-16T03:25:00Z') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokTokens|') + expect(wrapper.text()).not.toContain('2M|') + }) + + it.each([ + { tokens: 0, expected: 0, compact: '0' }, + { tokens: 1_000_000, expected: 50, compact: '1.0M' }, + { tokens: 2_000_000, expected: 100, compact: '2.0M' }, + { tokens: 2_200_000, expected: 100, compact: '2.2M' } + ])('Grok Free derives its 2M quota from local tokens: $tokens -> $expected%', async ({ tokens, expected, compact }) => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + plan: '' + }, + grok_local_usage_24h: { + requests: 5, + tokens, + cost: 0, + standard_cost: 0, + user_cost: 0 + }, + grok_request_quota: { limit: 100, remaining: 100 }, + grok_token_quota: { limit: 2_000_000, remaining: 2_000_000 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4300 + expected, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain(`24h|${expected}`) + expect(wrapper.findAll('span').filter((node) => node.text() === compact)).toHaveLength(1) + expect(wrapper.findAll('.usage-bar')).toHaveLength(1) + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokTokens|') + }) + + it('Grok Free uses rolling 24h usage instead of today-only usage', async () => { + getUsage.mockResolvedValue({ + grok_billing: { period_type: 'weekly', usage_percent: null, plan: '' }, + grok_local_usage: { + requests: 2, + tokens: 250_000, + cost: 0, + standard_cost: 0 + }, + grok_local_usage_24h: { + requests: 12, + tokens: 1_500_000, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4398, platform: 'grok', type: 'oauth', extra: {} }), + todayStats: { + requests: 2, + tokens: 200_000, + cost: 0, + standard_cost: 0 + } + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'title'], + template: '
{{ label }}|{{ utilization }}|{{ title }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('24h|75|admin.accounts.usageWindow.grokFreeQuota24hHint') + expect(wrapper.text()).toContain('1.5M') + expect(wrapper.text()).not.toContain('7d|') + expect(wrapper.text()).not.toContain('200.0K') + expect(wrapper.text()).not.toContain('250.0K') + }) + + it('Grok Free does not substitute today stats when rolling 24h usage is unavailable', async () => { + getUsage.mockResolvedValue({ + grok_billing: { period_type: 'weekly', usage_percent: null, plan: '' }, + grok_local_usage: { + requests: 1, + tokens: 250_000, + cost: 0, + standard_cost: 0, + user_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4399, platform: 'grok', type: 'oauth', extra: {} }), + todayStats: { + requests: 4, + tokens: 1_000_000, + cost: 0, + standard_cost: 0, + user_cost: 0 + } + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.findAll('.usage-bar')).toHaveLength(0) + expect(wrapper.text()).not.toContain('24h|') + expect(wrapper.text()).not.toContain('1.0M') + expect(wrapper.text()).not.toContain('250.0K') + }) + + it('Grok paid plans are not mistaken for Free when weekly usage is temporarily missing', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + plan: 'SuperGrok Heavy' + }, + grok_entitlement_status: 'free', + grok_local_usage: { + requests: 2, + tokens: 2_000_000, + cost: 1, + standard_cost: 1 + }, + grok_token_quota: { limit: 1_000, remaining: 250 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4401, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') + expect(wrapper.text()).not.toContain('2M|') + }) + + it('Grok custom paid monthly limits override stale Free entitlement', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + monthly_limit_cents: 25_000, + plan: '' + }, + grok_entitlement_status: 'free', + grok_local_usage: { + requests: 2, + tokens: 2_000_000, + cost: 1, + standard_cost: 1 + }, + grok_token_quota: { limit: 1_000, remaining: 250 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4402, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') + expect(wrapper.text()).not.toContain('2M|') + }) + + it('Grok credential Free tier keeps the 2M fallback when billing is unavailable', async () => { + getUsage.mockResolvedValue({ + subscription_tier: 'FREE', + grok_local_usage_24h: { + requests: 3, + tokens: 1_000_000, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4403, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('24h|50') + }) + + it('Grok paid manual probes keep the weekly/local summary when 24h usage is returned', async () => { + getUsage.mockResolvedValue({ + grok_quota_snapshot_state: 'no_headers', + error: 'stale error', + error_code: 'quota_unknown' + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4501, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt'], + template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: { + emits: ['probed'], + template: `` + } + } + } + }) + + await flushPromises() + await wrapper.get('.probe').trigger('click') + + expect(wrapper.text()).toContain('7d|42|2026-07-17T00:00:00Z') + expect(wrapper.text()).toContain('1.0M') + expect(wrapper.text()).not.toContain('750.0K') + expect(wrapper.text()).toContain('ACTIVE') + expect(wrapper.text()).not.toContain('stale error') + }) + + it('Grok Free manual probes merge rolling 24h usage', async () => { + getUsage.mockResolvedValue({ + subscription_tier: 'FREE', + grok_quota_snapshot_state: 'no_headers' + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4502, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: { + emits: ['probed'], + template: `` + } + } + } + }) + + await flushPromises() + await wrapper.get('.probe').trigger('click') + + expect(wrapper.text()).toContain('24h|75') + expect(wrapper.text()).toContain('1.5M') + expect(wrapper.text()).not.toContain('7d|') + }) + it('Key 账号在 today stats loading 时显示骨架屏', async () => { const wrapper = mount(AccountUsageCell, { props: { diff --git a/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts b/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts new file mode 100644 index 0000000000..0df3663f43 --- /dev/null +++ b/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts @@ -0,0 +1,19 @@ +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import { describe, expect, it } from 'vitest' + +const source = readFileSync( + resolve(process.cwd(), 'src/components/account/CreateAccountModal.vue'), + 'utf8' +) + +describe('CreateAccountModal Grok account types', () => { + it('offers API-key setup alongside OAuth with the official xAI default', () => { + expect(source).toContain('data-testid="grok-account-type-api-key"') + expect(source).toContain("@click=\"accountCategory = 'apikey'\"") + expect(source).toContain("newPlatform === 'grok'") + expect(source).toContain("? 'https://api.x.ai/v1'") + expect(source).toContain("form.platform === 'grok'") + expect(source).toContain("? 'xai-...'") + }) +}) diff --git a/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts new file mode 100644 index 0000000000..62c97d35a4 --- /dev/null +++ b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts @@ -0,0 +1,213 @@ +import { defineComponent } from 'vue' +import { flushPromises, mount } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { + createAccountMock, + importCodexSessionMock, + createOpenAICodexPATMock, +} = vi.hoisted(() => ({ + createAccountMock: vi.fn(), + importCodexSessionMock: vi.fn(), + createOpenAICodexPATMock: vi.fn(), +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => ({ + showError: vi.fn(), + showSuccess: vi.fn(), + showWarning: vi.fn(), + }), +})) + +vi.mock('@/stores/auth', () => ({ + useAuthStore: () => ({ isSimpleMode: true }), +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + accounts: { + create: createAccountMock, + checkMixedChannelRisk: vi.fn().mockResolvedValue({ has_risk: false }), + importCodexSession: importCodexSessionMock, + createOpenAICodexPAT: createOpenAICodexPATMock, + }, + settings: { + getWebSearchEmulationConfig: vi.fn().mockResolvedValue({ enabled: false, providers: [] }), + getSettings: vi.fn().mockResolvedValue({}), + }, + tlsFingerprintProfiles: { + list: vi.fn().mockResolvedValue([]), + }, + }, +})) + +vi.mock('@/api/admin/accounts', () => ({ + getAntigravityDefaultModelMapping: vi.fn().mockResolvedValue([]), +})) + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ t: (key: string) => key }), + } +}) + +import CreateAccountModal from '../CreateAccountModal.vue' + +const BaseDialogStub = defineComponent({ + name: 'BaseDialog', + props: { show: { type: Boolean, default: false } }, + template: '
', +}) + +const OAuthAuthorizationFlowStub = defineComponent({ + name: 'OAuthAuthorizationFlow', + emits: ['import-codex-session', 'import-codex-pat'], + template: ` +
+ + +
+ `, +}) + +function mountModal() { + return mount(CreateAccountModal, { + props: { show: true, proxies: [], groups: [] }, + global: { + stubs: { + BaseDialog: BaseDialogStub, + OAuthAuthorizationFlow: OAuthAuthorizationFlowStub, + ConfirmDialog: true, + Select: true, + Icon: true, + PlatformIcon: true, + ProxySelector: true, + ProxyAdBanner: true, + GroupSelector: true, + ModelWhitelistSelector: true, + QuotaLimitCard: true, + }, + }, + }) +} + +async function selectButtonByText(wrapper: ReturnType, text: string) { + const button = wrapper.findAll('button').find((candidate) => candidate.text().includes(text)) + expect(button).toBeDefined() + await button?.trigger('click') +} + +async function submitApiKeyAccount(platform: 'openai' | 'anthropic', enableLongContextBilling = false) { + const wrapper = mountModal() + await selectButtonByText(wrapper, platform === 'openai' ? 'OpenAI' : 'admin.accounts.claudeConsole') + if (platform === 'openai') { + await selectButtonByText(wrapper, 'API Key') + } + await wrapper.get('form#create-account-form input[type="text"]').setValue(`${platform} account`) + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + if (enableLongContextBilling) { + await wrapper.get('[data-testid="openai-long-context-billing-toggle"]').trigger('click') + } + await wrapper.get('form#create-account-form').trigger('submit.prevent') + await flushPromises() +} + +async function openCodexImportStep(toggleClicks = 0) { + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + for (let click = 0; click < toggleClicks; click += 1) { + await wrapper.get('[data-testid="openai-long-context-billing-toggle"]').trigger('click') + } + await wrapper.get('form#create-account-form input[type="text"]').setValue('Codex import') + await wrapper.get('form#create-account-form').trigger('submit.prevent') + return wrapper +} + +describe('CreateAccountModal OpenAI long-context billing', () => { + beforeEach(() => { + createAccountMock.mockReset().mockResolvedValue({}) + importCodexSessionMock.mockReset().mockResolvedValue({ + created: 1, + updated: 0, + skipped: 0, + failed: 0, + errors: [], + warnings: [], + }) + createOpenAICodexPATMock.mockReset().mockResolvedValue({}) + }) + + it('sends false explicitly for normal OpenAI account creation by default', async () => { + await submitApiKeyAccount('openai') + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('sends true explicitly when OpenAI long-context billing is enabled', async () => { + await submitApiKeyAccount('openai', true) + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('omits the OpenAI setting for non-OpenAI account creation', async () => { + await submitApiKeyAccount('anthropic') + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('leaves Codex session import billing ownership to the backend', async () => { + const wrapper = await openCodexImportStep() + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock).toHaveBeenCalledTimes(1) + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('leaves Codex PAT import billing ownership to the backend', async () => { + const wrapper = await openCodexImportStep() + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock).toHaveBeenCalledTimes(1) + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('sends explicit true for Codex session import after the toggle is enabled', async () => { + const wrapper = await openCodexImportStep(1) + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('sends explicit false for Codex session import after the toggle is changed back', async () => { + const wrapper = await openCodexImportStep(2) + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('sends explicit true for Codex PAT import after the toggle is enabled', async () => { + const wrapper = await openCodexImportStep(1) + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('sends explicit false for Codex PAT import after the toggle is changed back', async () => { + const wrapper = await openCodexImportStep(2) + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) +}) diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index c148d04a6f..b3a583d102 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -267,6 +267,18 @@ function buildGrokOAuthAccount() { } as any } +function buildGrokAPIKeyAccount() { + return { + ...buildAccount(), + id: 6, + name: 'Grok API Key', + platform: 'grok', + credentials: {}, + credentials_status: { has_api_key: true }, + concurrency: 2 + } as any +} + function buildOpenAISetupTokenAccount() { return { ...buildAccount(), @@ -383,6 +395,105 @@ describe('EditAccountModal', () => { }) }) + it('loads and submits the per-account OpenAI long-context billing toggle', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: true + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('true') + + await toggle.trigger('click') + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('defaults legacy OpenAI accounts to long-context billing disabled', async () => { + const account = buildAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('does not render or submit the long-context billing toggle for Spark shadow accounts', async () => { + const account = buildOpenAISparkShadowAccount() + account.extra = { + openai_long_context_billing_enabled: false + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + const wrapper = mountModal(account) + + expect(wrapper.find('[data-testid="openai-long-context-billing-toggle"]').exists()).toBe(false) + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra).not.toHaveProperty( + 'openai_long_context_billing_enabled' + ) + }) + + it('preserves an explicit OpenAI long-context billing opt-out', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: false + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('fails closed for malformed OpenAI long-context billing values', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: 'false' + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + expect(wrapper.get('[data-testid="openai-long-context-billing-toggle"]').attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + it('loads and submits Grok OAuth model mapping edits', async () => { const account = buildGrokOAuthAccount() updateAccountMock.mockReset() @@ -412,6 +523,24 @@ describe('EditAccountModal', () => { }) }) + it('uses the official xAI base URL when a Grok API-key account omits base_url', async () => { + const account = buildGrokAPIKeyAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + expect((wrapper.get('input[placeholder="https://api.x.ai/v1"]').element as HTMLInputElement).value) + .toBe('https://api.x.ai/v1') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.base_url).toBe('https://api.x.ai/v1') + }) + it('only submits model mapping credentials when saving an OpenAI spark shadow account', async () => { authIsSimpleMode.value = false const account = buildOpenAISparkShadowAccount() diff --git a/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts b/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts new file mode 100644 index 0000000000..9b431b339e --- /dev/null +++ b/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts @@ -0,0 +1,54 @@ +import { flushPromises, mount } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' +import GrokQuotaProbeCell from '../GrokQuotaProbeCell.vue' +import type { Account } from '@/types' + +const { queryQuota } = vi.hoisted(() => ({ + queryQuota: vi.fn() +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + grok: { queryQuota } + } +})) + +vi.mock('vue-i18n', () => ({ + useI18n: () => ({ + t: (key: string, params?: Record) => + params?.percent == null ? key : `${key}:${params.percent}` + }) +})) + +const account = { + id: 99, + platform: 'grok', + type: 'oauth' +} as Account + +describe('GrokQuotaProbeCell', () => { + beforeEach(() => { + queryQuota.mockReset() + }) + + it('keeps billing data while exposing a failed Free quota fallback', async () => { + queryQuota.mockResolvedValue({ + source: 'hybrid_probe', + billing: { period_type: 'weekly', usage_percent: null }, + headers_observed: false, + reset_supported: false, + fetched_at: 1, + probe_error: 'upstream returned 402 for probe model "grok-4.5"' + }) + const wrapper = mount(GrokQuotaProbeCell, { props: { account } }) + + await wrapper.get('button').trigger('click') + await flushPromises() + + expect(wrapper.text()).toContain('upstream returned 402 for probe model "grok-4.5"') + expect(wrapper.emitted('probed')?.[0]?.[0]).toMatchObject({ + billing: { period_type: 'weekly', usage_percent: null }, + probe_error: 'upstream returned 402 for probe model "grok-4.5"' + }) + }) +}) diff --git a/frontend/src/components/account/__tests__/UsageProgressBar.spec.ts b/frontend/src/components/account/__tests__/UsageProgressBar.spec.ts index 6fa6575f54..af5fc5d66d 100644 --- a/frontend/src/components/account/__tests__/UsageProgressBar.spec.ts +++ b/frontend/src/components/account/__tests__/UsageProgressBar.spec.ts @@ -96,4 +96,54 @@ describe('UsageProgressBar', () => { expect(wrapper.text()).toContain('usage.resetNow') expect(wrapper.text()).not.toContain('usage.resetPending') }) + + it('剩余容量模式在 100% 时显示满格绿色', () => { + const wrapper = mount(UsageProgressBar, { + props: { + label: 'Req', + utilization: 100, + remainingCapacity: true, + color: 'indigo' + } + }) + + expect(wrapper.text()).toContain('100%') + expect(wrapper.get('.h-1\\.5 > div').attributes('style')).toContain('width: 100%') + expect(wrapper.get('.h-1\\.5 > div').classes()).toContain('bg-green-500') + }) + + it('剩余容量模式在低量和耗尽时缩短并变红', async () => { + const wrapper = mount(UsageProgressBar, { + props: { + label: 'Req', + utilization: 15, + remainingCapacity: true, + color: 'indigo' + } + }) + + expect(wrapper.text()).toContain('15%') + expect(wrapper.get('.h-1\\.5 > div').attributes('style')).toContain('width: 15%') + expect(wrapper.get('.h-1\\.5 > div').classes()).toContain('bg-red-500') + + await wrapper.setProps({ utilization: 0 }) + + expect(wrapper.text()).toContain('0%') + expect(wrapper.get('.h-1\\.5 > div').attributes('style')).toContain('width: 0%') + expect(wrapper.get('.h-1\\.5 > div').classes()).toContain('bg-red-500') + }) + + it('默认利用率模式仍把超限显示为满格红色', () => { + const wrapper = mount(UsageProgressBar, { + props: { + label: '5h', + utilization: 120, + color: 'indigo' + } + }) + + expect(wrapper.text()).toContain('120%') + expect(wrapper.get('.h-1\\.5 > div').attributes('style')).toContain('width: 100%') + expect(wrapper.get('.h-1\\.5 > div').classes()).toContain('bg-red-500') + }) }) diff --git a/frontend/src/components/account/__tests__/credentialsBuilder.spec.ts b/frontend/src/components/account/__tests__/credentialsBuilder.spec.ts index c2cb093805..ae4c6739a8 100644 --- a/frontend/src/components/account/__tests__/credentialsBuilder.spec.ts +++ b/frontend/src/components/account/__tests__/credentialsBuilder.spec.ts @@ -6,9 +6,13 @@ import { applyAntigravityProjectID, applyHeaderOverride, applyInterceptWarmup, + applyPlanType, buildHeaderOverridesObject, + buildPlanTypeOptions, getHeaderOverrideTemplate, isHeaderOverridePlatform, + planTypeDisplayLabel, + readPlanType, splitHeaderOverridesObject, validateHeaderOverrideRows } from '../credentialsBuilder' @@ -289,3 +293,88 @@ describe('validateHeaderOverrideRows session isolation headers', () => { expect(validateHeaderOverrideRows([{ name: 'x'.repeat(201), value: 'v' }])).toBe('invalidName') }) }) + +describe('plan_type helpers', () => { + describe('planTypeDisplayLabel', () => { + it('maps canonical + alias values to friendly labels', () => { + expect(planTypeDisplayLabel('plus')).toBe('Plus') + expect(planTypeDisplayLabel('pro')).toBe('Pro') + expect(planTypeDisplayLabel('chatgptpro')).toBe('Pro') + expect(planTypeDisplayLabel('free')).toBe('Free') + expect(planTypeDisplayLabel('team')).toBe('Team') + expect(planTypeDisplayLabel('CHATGPTPRO')).toBe('Pro') + }) + it('returns unknown values verbatim', () => { + expect(planTypeDisplayLabel('self_serve_business')).toBe('self_serve_business') + }) + }) + + describe('readPlanType', () => { + it('reads a string plan_type', () => { + expect(readPlanType({ plan_type: 'plus' })).toBe('plus') + }) + it('treats non-string / missing values as empty', () => { + expect(readPlanType({ plan_type: 42 })).toBe('') + expect(readPlanType({ plan_type: true })).toBe('') + expect(readPlanType({})).toBe('') + expect(readPlanType(undefined)).toBe('') + expect(readPlanType(null)).toBe('') + }) + }) + + describe('buildPlanTypeOptions', () => { + const clear = 'Clear' + it('returns clear + presets when current is empty', () => { + expect(buildPlanTypeOptions('', clear)).toEqual([ + { value: '', label: clear }, + { value: 'plus', label: 'Plus' }, + { value: 'pro', label: 'Pro' }, + { value: 'free', label: 'Free' } + ]) + }) + it('keeps canonical chatgptpro under a single friendly "Pro" option (no duplicate)', () => { + const opts = buildPlanTypeOptions('chatgptpro', clear) + const pros = opts.filter(o => o.label === 'Pro') + expect(pros).toHaveLength(1) + expect(pros[0].value).toBe('chatgptpro') + expect(opts.map(o => o.value)).toEqual(['', 'plus', 'chatgptpro', 'free']) + }) + it('appends an unknown-but-labeled value (team) as its own option', () => { + const opts = buildPlanTypeOptions('team', clear) + expect(opts.find(o => o.value === 'team')).toEqual({ value: 'team', label: 'Team' }) + // presets untouched + expect(opts.map(o => o.value)).toEqual(['', 'plus', 'pro', 'free', 'team']) + }) + it('appends a fully custom value with a raw label', () => { + const opts = buildPlanTypeOptions('weird_x', clear) + expect(opts.at(-1)).toEqual({ value: 'weird_x', label: 'weird_x' }) + }) + it('does not duplicate an exact preset value', () => { + const opts = buildPlanTypeOptions('pro', clear) + expect(opts.filter(o => o.value === 'pro')).toHaveLength(1) + expect(opts.map(o => o.value)).toEqual(['', 'plus', 'pro', 'free']) + }) + }) + + describe('applyPlanType', () => { + it('sets plan_type and preserves all other credential keys', () => { + const creds = { + chatgpt_account_id: 'acc', + email: 'a@b.c', + subscription_expires_at: '2026-01-01', + model_mapping: { x: 'y' } + } + const out = applyPlanType({ ...creds }, 'plus') + expect(out).toEqual({ ...creds, plan_type: 'plus' }) + }) + it('trims the value', () => { + expect(applyPlanType({}, ' pro ')).toEqual({ plan_type: 'pro' }) + }) + it('deletes the key when cleared (empty), keeping other keys', () => { + const out = applyPlanType({ plan_type: 'pro', email: 'a@b.c' }, '') + expect(out).toEqual({ email: 'a@b.c' }) + expect('plan_type' in out).toBe(false) + }) + }) +}) + diff --git a/frontend/src/components/account/credentialsBuilder.ts b/frontend/src/components/account/credentialsBuilder.ts index 3cdc0bdd54..e78cf41de6 100644 --- a/frontend/src/components/account/credentialsBuilder.ts +++ b/frontend/src/components/account/credentialsBuilder.ts @@ -201,3 +201,87 @@ export function applyHeaderOverride( delete credentials[HEADER_OVERRIDES_CREDENTIAL_KEY] } } + +// ===== OpenAI plan_type (ChatGPT 订阅档位) 手动覆盖 ===== + +export interface PlanTypeOption { + value: string + label: string + // 兼容 common/Select.vue 的 SelectOption(含索引签名) + [key: string]: unknown +} + +/** + * plan_type 值的友好显示标签,镜像 PlatformTypeBadge 的映射 + * (canonical 值 chatgptpro 显示为 Pro,team 显示为 Team)。未知值原样返回。 + */ +export function planTypeDisplayLabel(value: string): string { + switch (value.trim().toLowerCase()) { + case 'plus': + return 'Plus' + case 'pro': + case 'chatgptpro': + return 'Pro' + case 'free': + return 'Free' + case 'team': + return 'Team' + default: + return value + } +} + +/** + * 从凭据里读取 plan_type,仅接受字符串(脏数据 42/true 等一律视为空, + * 避免被当作合法自定义项保留)。 + */ +export function readPlanType(credentials: Record | undefined | null): string { + const v = credentials?.plan_type + return typeof v === 'string' ? v : '' +} + +/** + * 构建 plan_type 下拉选项:清空 + Plus/Pro/Free 预设。 + * 若当前值是某预设的别名(如 chatgptpro↔Pro),用当前的 canonical 值占据该 + * 标签位(保留 canonical,显示友好标签,避免重复项);若是完全预设外的值 + * (如 team 或异常值),追加为一项,避免编辑时下拉丢失原值。 + */ +export function buildPlanTypeOptions(current: string, clearLabel: string): PlanTypeOption[] { + const cur = (current || '').trim() + const curLabel = cur ? planTypeDisplayLabel(cur) : '' + const presets: PlanTypeOption[] = [ + { value: 'plus', label: 'Plus' }, + { value: 'pro', label: 'Pro' }, + { value: 'free', label: 'Free' } + ] + const opts: PlanTypeOption[] = [{ value: '', label: clearLabel }] + for (const p of presets) { + if (cur && p.value !== cur.toLowerCase() && p.label === curLabel) { + // 当前值是该预设的别名:用 canonical 当前值占位,标签仍显示友好名 + opts.push({ value: cur, label: p.label }) + } else { + opts.push(p) + } + } + if (cur && !opts.some(o => o.value.toLowerCase() === cur.toLowerCase())) { + opts.push({ value: cur, label: planTypeDisplayLabel(cur) }) + } + return opts +} + +/** + * 把手动选择的 plan_type 写入凭据:非空则设置,空则删除该键(清空/自动识别)。 + * 直接修改传入对象并返回。 + */ +export function applyPlanType( + credentials: Record, + planType: string +): Record { + const pt = (planType || '').trim() + if (pt) { + credentials.plan_type = pt + } else { + delete credentials.plan_type + } + return credentials +} diff --git a/frontend/src/components/admin/account/AccountTestModal.vue b/frontend/src/components/admin/account/AccountTestModal.vue index 0a0e3dd9ae..0a8f853ebb 100644 --- a/frontend/src/components/admin/account/AccountTestModal.vue +++ b/frontend/src/components/admin/account/AccountTestModal.vue @@ -250,6 +250,7 @@ import TextArea from '@/components/common/TextArea.vue' import { Icon } from '@/components/icons' import { useClipboard } from '@/composables/useClipboard' import { buildApiUrl } from '@/api/client' +import { ADMIN_UI_REQUEST_HEADER } from '@/api/adminUIRequest' import { adminAPI } from '@/api/admin' import type { Account, ClaudeModel } from '@/types' @@ -438,7 +439,8 @@ const startTest = async () => { method: 'POST', headers: { Authorization: `Bearer ${localStorage.getItem('auth_token')}`, - 'Content-Type': 'application/json' + 'Content-Type': 'application/json', + [ADMIN_UI_REQUEST_HEADER]: '1' }, body: JSON.stringify(requestBody), signal: abortController.signal diff --git a/frontend/src/components/admin/monitor/MonitorAdvancedRequestConfig.vue b/frontend/src/components/admin/monitor/MonitorAdvancedRequestConfig.vue index 404b691692..c4ccdfa3ac 100644 --- a/frontend/src/components/admin/monitor/MonitorAdvancedRequestConfig.vue +++ b/frontend/src/components/admin/monitor/MonitorAdvancedRequestConfig.vue @@ -109,6 +109,8 @@ import { useI18n } from 'vue-i18n' import type { APIMode, BodyOverrideMode, Provider } from '@/api/admin/channelMonitor' import { API_MODE_RESPONSES, + DEFAULT_GROK_MODEL, + PROVIDER_GROK, PROVIDER_OPENAI, } from '@/constants/channelMonitor' @@ -305,11 +307,12 @@ const bodyPlaceholder = computed(() => { } return '{\n "model": "gpt-4o-mini",\n "instructions": "You are a health check endpoint. Reply briefly.",\n "input": "Reply with exactly: ok",\n "max_output_tokens": 20,\n "stream": false\n}' } - if (props.provider === PROVIDER_OPENAI) { + if (props.provider === PROVIDER_OPENAI || props.provider === PROVIDER_GROK) { if (props.bodyOverrideMode === 'merge') { return '{\n "max_tokens": 20\n}' } - return '{\n "model": "gpt-4o-mini",\n "messages": [{"role":"user","content":"Reply with exactly: ok"}],\n "max_tokens": 20,\n "stream": false\n}' + const model = props.provider === PROVIDER_GROK ? DEFAULT_GROK_MODEL : 'gpt-4o-mini' + return `{\n "model": "${model}",\n "messages": [{"role":"user","content":"Reply with exactly: ok"}],\n "max_tokens": 20,\n "stream": false\n}` } if (props.bodyOverrideMode === 'merge') { return '{\n "system": "You are Claude Code..."\n}' diff --git a/frontend/src/components/admin/monitor/MonitorFiltersBar.vue b/frontend/src/components/admin/monitor/MonitorFiltersBar.vue index eb2a5c7857..544238f49c 100644 --- a/frontend/src/components/admin/monitor/MonitorFiltersBar.vue +++ b/frontend/src/components/admin/monitor/MonitorFiltersBar.vue @@ -70,6 +70,7 @@ import { PROVIDER_OPENAI, PROVIDER_ANTHROPIC, PROVIDER_GEMINI, + PROVIDER_GROK, } from '@/constants/channelMonitor' defineProps<{ @@ -94,6 +95,7 @@ const providerFilterOptions = computed(() => [ { value: PROVIDER_OPENAI, label: t('monitorCommon.providers.openai') }, { value: PROVIDER_ANTHROPIC, label: t('monitorCommon.providers.anthropic') }, { value: PROVIDER_GEMINI, label: t('monitorCommon.providers.gemini') }, + { value: PROVIDER_GROK, label: t('monitorCommon.providers.grok') }, ]) const enabledFilterOptions = computed(() => [ diff --git a/frontend/src/components/admin/monitor/MonitorFormDialog.vue b/frontend/src/components/admin/monitor/MonitorFormDialog.vue index e6cab8edf1..14a9e2dd15 100644 --- a/frontend/src/components/admin/monitor/MonitorFormDialog.vue +++ b/frontend/src/components/admin/monitor/MonitorFormDialog.vue @@ -13,15 +13,16 @@
-
+
@@ -80,6 +81,7 @@ (() => [ { value: PROVIDER_ANTHROPIC, label: t('monitorCommon.providers.anthropic') }, { value: PROVIDER_OPENAI, label: t('monitorCommon.providers.openai') }, { value: PROVIDER_GEMINI, label: t('monitorCommon.providers.gemini') }, + { value: PROVIDER_GROK, label: t('monitorCommon.providers.grok') }, ]) +function selectProvider(provider: Provider) { + if (form.provider === provider) return + const previousProvider = form.provider + const clearGrokEndpoint = + previousProvider === PROVIDER_GROK && form.endpoint === DEFAULT_GROK_ENDPOINT + const clearGrokModel = + previousProvider === PROVIDER_GROK && form.primary_model === DEFAULT_GROK_MODEL + form.provider = provider + if (provider === PROVIDER_GROK) { + if (!form.endpoint.trim()) form.endpoint = DEFAULT_GROK_ENDPOINT + if (!form.primary_model.trim()) form.primary_model = DEFAULT_GROK_MODEL + return + } + if (clearGrokEndpoint) form.endpoint = '' + if (clearGrokModel) form.primary_model = '' +} + // Clear api_key whenever provider changes to avoid cross-provider key mismatch. // Editing mode loads api_key='' via loadFromMonitor and only sets it on user // typing, so clearing on provider change is always a safe no-op until the user diff --git a/frontend/src/components/admin/monitor/MonitorTemplateManagerDialog.vue b/frontend/src/components/admin/monitor/MonitorTemplateManagerDialog.vue index e54ecf631d..63b87db75a 100644 --- a/frontend/src/components/admin/monitor/MonitorTemplateManagerDialog.vue +++ b/frontend/src/components/admin/monitor/MonitorTemplateManagerDialog.vue @@ -7,7 +7,7 @@ >
-
+