diff --git a/README.md b/README.md index 16ebab0a32..5d0c4f433a 100644 --- a/README.md +++ b/README.md @@ -58,6 +58,11 @@ Please read the following carefully before using this project: Thanks to AIGoCode for sponsoring this project! AIGoCode is an all-in-one platform that integrates Claude Code, Codex, and the latest Gemini models, providing you with stable, efficient, and highly cost-effective AI coding services. The platform offers flexible subscription plans, zero risk of account suspension, direct access with no VPN required, and lightning-fast responses. AIGoCode has prepared a special benefit for sub2api users: if you register via this link, you'll receive an extra 10% bonus credit on your first top-up! + +CodexEverywhere +Real GPT-5.6 series at 3% of OpenAI pricing — CodexEverywhere is democratizing access to frontier models for developers worldwide. We believe in transparency and honesty, with model quality verified by active community oversight for months. USD and crypto friendly. Start with a free $20 trial at codex-everywhere.com. + + bmoplus Huge thanks to BmoPlus for sponsoring this project! BmoPlus is a highly reliable AI account provider built strictly for heavy AI users and developers. They offer rock-solid, ready-to-use accounts and official top-up services for ChatGPT Plus / ChatGPT Pro (Full Warranty) / Claude Pro / Super Grok / Gemini Pro. By registering and ordering through BmoPlus - Premium AI Accounts & Top-ups, users can unlock the mind-blowing rate of 10% of the official GPT subscription price (90% OFF) diff --git a/README_CN.md b/README_CN.md index e879045374..faa0200c58 100644 --- a/README_CN.md +++ b/README_CN.md @@ -59,6 +59,11 @@ 感谢 AIGoCode 赞助了本项目!AIGoCode 是一站式集成 Claude Code、Codex 以及最新 Gemini 模型的综合平台,为您提供稳定、高效、高性价比的 AI 编程服务。平台提供灵活的订阅方案,零封号风险,免 VPN 直连,响应极速。AIGoCode 为 sub2api 用户准备了专属福利:通过此链接注册,首次充值可额外获得 10% 赠送额度! + +CodexEverywhere +Real GPT-5.6 series at 3% of OpenAI pricing — CodexEverywhere is democratizing access to frontier models for developers worldwide. We believe in transparency and honesty, with model quality verified by active community oversight for months. USD and crypto friendly. Start with a free $20 trial at codex-everywhere.com. + + bmoplus 感谢 BmoPlus 赞助了本项目!BmoPlus 是一家专为AI订阅重度用户打造的可靠 AI 账号代充服务商,提供稳定的 ChatGPT Plus / ChatGPT Pro(全程质保) / Claude Pro / Super Grok / Gemini Pro 的官方代充&成品账号。 通过BmoPlus AI成品号专卖/代充注册下单的用户,可享GPT 官网订阅一折 的震撼价格! diff --git a/README_JA.md b/README_JA.md index 17890e4557..9bd3802a32 100644 --- a/README_JA.md +++ b/README_JA.md @@ -58,6 +58,11 @@ AIGoCode のご支援に感謝します!AIGoCode は Claude Code、Codex、最新の Gemini モデルを統合したオールインワンプラットフォームで、安定的かつ効率的でコストパフォーマンスに優れた AI コーディングサービスを提供します。柔軟なサブスクリプションプラン、アカウント停止リスクゼロ、VPN 不要の直接アクセス、超高速レスポンスが特長です。AIGoCode は sub2api ユーザー向けに特別特典を用意しています:こちらのリンクから登録すると、初回チャージ時に 10% のボーナスクレジットを追加プレゼント! + +CodexEverywhere +OpenAI 公式価格のわずか 3% で本物の GPT-5.6 シリーズを提供 — CodexEverywhere は世界中の開発者にフロンティアモデルへのアクセスを民主化しています。私たちは透明性と誠実さを信条とし、モデル品質は数か月にわたるアクティブなコミュニティの監視によって検証されています。USD および暗号通貨に対応。codex-everywhere.com で $20 の無料トライアルから始めましょう。 + + bmoplus 本プロジェクトにご支援いただいた BmoPlus に感謝いたします!BmoPlusは、AIサブスクリプションのヘビーユーザー向けに特化した信頼性の高いAIアカウントサービスプロバイダーであり、安定した ChatGPT Plus / ChatGPT Pro (完全保証) / Claude Pro / Super Grok / Gemini Pro の公式代行チャージおよび即納アカウントを提供しています。こちらのBmoPlus AIアカウント専門店/代行チャージ経由でご登録・ご注文いただいたユーザー様は、GPTを 公式サイト価格の約1割(90% OFF) という驚異的な価格でご利用いただけます! diff --git a/assets/partners/logos/codex-everywhere.jpg b/assets/partners/logos/codex-everywhere.jpg new file mode 100644 index 0000000000..a89679ada5 Binary files /dev/null and b/assets/partners/logos/codex-everywhere.jpg differ diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index ce08f35612..24f1cf39c7 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.182 +0.1.183 diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 93067da278..599a04a34e 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -2774,13 +2774,15 @@ func (h *AccountHandler) SyncUpstreamModels(c *gin.Context) { return } - models, err := h.accountTestService.FetchUpstreamSupportedModels(c.Request.Context(), account) + catalog, err := h.accountTestService.SyncUpstreamModelCatalog(c.Request.Context(), account) if err != nil { var syncErr *service.UpstreamModelSyncError if errors.As(err, &syncErr) { switch syncErr.Kind { case service.UpstreamModelSyncErrorConfiguration, service.UpstreamModelSyncErrorUnsupported: response.BadRequest(c, syncErr.SafeMessage()) + case service.UpstreamModelSyncErrorInternal: + response.InternalError(c, syncErr.SafeMessage()) default: slog.Warn("sync_upstream_models_failed", "account_id", accountID, "kind", syncErr.Kind) response.Error(c, http.StatusBadGateway, syncErr.SafeMessage()) @@ -2793,29 +2795,35 @@ func (h *AccountHandler) SyncUpstreamModels(c *gin.Context) { return } - response.Success(c, gin.H{"models": models}) + response.Success(c, catalog) } // SyncUpstreamModelsPreview handles syncing live supported models using provided credentials (no account ID needed). // POST /api/v1/admin/accounts/models/sync-upstream-preview func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) { var req struct { - Platform string `json:"platform" binding:"required"` - Type string `json:"type" binding:"required"` - BaseURL string `json:"base_url"` - APIKey string `json:"api_key" binding:"required"` + Platform string `json:"platform" binding:"required"` + Type string `json:"type" binding:"required"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key" binding:"required"` + ModelMapping map[string]string `json:"model_mapping"` } if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } + modelMapping := make(map[string]any, len(req.ModelMapping)) + for sourceModel, upstreamModel := range req.ModelMapping { + modelMapping[sourceModel] = upstreamModel + } tempAccount := &service.Account{ Platform: req.Platform, Type: req.Type, Credentials: map[string]any{ - "api_key": req.APIKey, - "base_url": req.BaseURL, + "api_key": req.APIKey, + "base_url": req.BaseURL, + "model_mapping": modelMapping, }, } @@ -2824,13 +2832,15 @@ func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) { return } - models, err := h.accountTestService.FetchUpstreamSupportedModels(c.Request.Context(), tempAccount) + catalog, err := h.accountTestService.SyncUpstreamModelCatalog(c.Request.Context(), tempAccount) if err != nil { var syncErr *service.UpstreamModelSyncError if errors.As(err, &syncErr) { switch syncErr.Kind { case service.UpstreamModelSyncErrorConfiguration, service.UpstreamModelSyncErrorUnsupported: response.BadRequest(c, syncErr.SafeMessage()) + case service.UpstreamModelSyncErrorInternal: + response.InternalError(c, syncErr.SafeMessage()) default: slog.Warn("sync_upstream_models_preview_failed", "platform", req.Platform, "kind", syncErr.Kind) response.Error(c, http.StatusBadGateway, syncErr.SafeMessage()) @@ -2843,7 +2853,7 @@ func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) { return } - response.Success(c, gin.H{"models": models}) + response.Success(c, catalog) } // SetPrivacy handles setting privacy for a single OpenAI/Antigravity OAuth account diff --git a/backend/internal/handler/admin/account_handler_available_models_test.go b/backend/internal/handler/admin/account_handler_available_models_test.go index ef2f65d4e3..28dc6d1492 100644 --- a/backend/internal/handler/admin/account_handler_available_models_test.go +++ b/backend/internal/handler/admin/account_handler_available_models_test.go @@ -38,14 +38,20 @@ func setupAvailableModelsRouter(adminSvc service.AdminService) *gin.Engine { } type syncUpstreamHTTPUpstream struct { - resp *http.Response - err error + resp *http.Response + responses []*http.Response + err error } func (u *syncUpstreamHTTPUpstream) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { if u.err != nil { return nil, u.err } + if len(u.responses) > 0 { + resp := u.responses[0] + u.responses = u.responses[1:] + return resp, nil + } return u.resp, nil } @@ -68,6 +74,7 @@ func setupSyncUpstreamModelsRouter(adminSvc service.AdminService, upstream servi ) handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, accountTestSvc, nil, nil, nil, nil, nil) router.POST("/api/v1/admin/accounts/:id/models/sync-upstream", handler.SyncUpstreamModels) + router.POST("/api/v1/admin/accounts/models/sync-upstream-preview", handler.SyncUpstreamModelsPreview) return router } @@ -347,6 +354,99 @@ func TestAccountHandlerSyncUpstreamModels_ConfigErrorReturnsBadRequest(t *testin require.Contains(t, rec.Body.String(), "No OpenAI API key is available") } +func TestAccountHandlerSyncUpstreamModelsReturnsCapabilityMetadata(t *testing.T) { + svc := &availableModelsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 48, Name: "custom-openai", Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, Status: service.StatusActive, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + }, + } + upstream := &syncUpstreamHTTPUpstream{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"models":[{ + "id":"custom-thinking-model", + "reasoning":true, + "default_reasoning_level":"high", + "supported_reasoning_levels":["low","high"], + "input_modalities":["text","image"], + "context_window":256000 + }]}`)), + }} + router := setupSyncUpstreamModelsRouter(svc, upstream) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/48/models/sync-upstream", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var resp struct { + Data service.UpstreamModelCatalog `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, []string{"custom-thinking-model"}, resp.Data.Models) + metadata := resp.Data.Metadata["custom-thinking-model"] + require.NotNil(t, metadata.Reasoning) + require.True(t, *metadata.Reasoning) + require.Equal(t, []string{"low", "high"}, metadata.SupportedReasoningLevels) + require.Equal(t, []string{"text", "image"}, metadata.InputModalities) +} + +// Scenario: 创建账号 preview 将具体 mapping 传给 404/405 配置回退。 +func TestAccountHandlerSyncUpstreamModelsPreviewUsesProvidedModelMapping(t *testing.T) { + upstream := &syncUpstreamHTTPUpstream{responses: []*http.Response{ + { + StatusCode: http.StatusNotFound, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":"not found"}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "configured-provider": { + "api": "https://provider.example/v1", + "models": { + "glm-5.3": { + "id": "glm-5.3", + "reasoning": true, + "reasoning_options": [{"type":"effort","values":["low","high"]}], + "modalities": {"input":["text"],"output":["text"]}, + "limit": {"context":1000000,"output":131072} + } + } + } + }`)), + }, + }} + router := setupSyncUpstreamModelsRouter(newStubAdminService(), upstream) + + rec := httptest.NewRecorder() + req := httptest.NewRequest( + http.MethodPost, + "/api/v1/admin/accounts/models/sync-upstream-preview", + strings.NewReader(`{ + "platform":"openai", + "type":"apikey", + "base_url":"https://provider.example/v1", + "api_key":"key", + "model_mapping":{"public-glm":"glm-5.3"} + }`), + ) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var resp struct { + Data service.UpstreamModelCatalog `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, []string{"glm-5.3"}, resp.Data.Models) + require.Equal(t, []string{"low", "high"}, resp.Data.Metadata["glm-5.3"].SupportedReasoningLevels) +} + func TestAccountHandlerSyncUpstreamModels_UpstreamErrorDoesNotExposeBody(t *testing.T) { svc := &availableModelsAdminService{ stubAdminService: newStubAdminService(), @@ -377,3 +477,52 @@ func TestAccountHandlerSyncUpstreamModels_UpstreamErrorDoesNotExposeBody(t *test require.Contains(t, rec.Body.String(), "Upstream model list request failed with HTTP 502") require.NotContains(t, rec.Body.String(), "SECRET_TOKEN") } + +// Scenario: 能力补全失败显示部分成功。 +func TestAccountHandlerSyncUpstreamModels_MetadataEnrichmentFailureReturnsWarning(t *testing.T) { + svc := &availableModelsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 46, + Name: "opencode-id-only-model-list", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Credentials: map[string]any{ + "api_key": "opencode-key", + "base_url": "https://opencode.ai/zen/v1", + }, + }, + } + upstream := &syncUpstreamHTTPUpstream{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"x-preview-f-free"}]}`)), + }, + { + StatusCode: http.StatusBadGateway, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":"registry unavailable"}`)), + }, + }} + router := setupSyncUpstreamModelsRouter(svc, upstream) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/46/models/sync-upstream", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + var resp struct { + Data struct { + Models []string `json:"models"` + Warnings []struct { + Code string `json:"code"` + } `json:"warnings"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, []string{"x-preview-f-free"}, resp.Data.Models) + require.Len(t, resp.Data.Warnings, 1) + require.Equal(t, "upstream_model_metadata_incomplete", resp.Data.Warnings[0].Code) +} diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 6985392926..d13bbee47a 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -312,6 +312,14 @@ func GetInboundEndpoint(c *gin.Context) string { // and the account platform. Handlers call this after scheduling an // account, passing account.Platform. func GetUpstreamEndpoint(c *gin.Context, platform string) string { + // OpenAI 转发服务维护独立的运行时端点上下文,覆盖普通入站推导。 + // 这对 force_chat_completions 的错误路径尤为重要:此时可能没有 + // ForwardResult,不能把入站 /v1/responses 误报成上游端点。 + if platform == service.PlatformOpenAI || platform == service.PlatformGrok || service.IsCNProvider(platform) { + if endpoint := service.GetActualOpenAIUpstreamEndpoint(c); endpoint != "" { + return endpoint + } + } if c != nil { if value, ok := c.Get(ctxKeyActualUpstreamEndpoint); ok { if endpoint, ok := value.(string); ok && endpoint != "" { diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index 005c397d44..f87b27fcd1 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -184,6 +184,16 @@ func TestGetUpstreamEndpointPrefersRuntimeOverride(t *testing.T) { require.Equal(t, EndpointMessages, GetUpstreamEndpoint(c, service.PlatformAntigravity)) } +func TestGetUpstreamEndpointUsesOpenAIRuntimeOverride(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil) + c.Set(ctxKeyInboundEndpoint, EndpointResponses) + + service.SetActualOpenAIUpstreamEndpoint(c, EndpointChatCompletions) + require.Equal(t, EndpointChatCompletions, GetUpstreamEndpoint(c, service.PlatformOpenAI)) +} + func TestResolveOpenAIUpstreamEndpointPrefersForwardResult(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 1128329be8..d07cd57e13 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -1140,6 +1140,79 @@ func (h *GatewayHandler) Models(c *gin.Context) { }) } +// CodexModels returns the effective group model list using the manifest shape +// expected by Codex custom providers. Official OpenAI groups continue to use +// OpenAIGatewayHandler.CodexModels so their live upstream metadata is preserved. +func (h *GatewayHandler) CodexModels(c *gin.Context) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil || apiKey.Group == nil { + h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required") + return + } + + forcedPlatform := "" + if value, exists := middleware2.GetForcePlatformFromContext(c); exists { + forcedPlatform = strings.TrimSpace(value) + } + modelIDs := h.codexModelIDsForGroup(c.Request.Context(), apiKey.Group, forcedPlatform) + modelIDs = service.FilterCodexModelIDsForGroup(modelIDs, apiKey.Group) + body, err := h.gatewayService.BuildCodexModelsManifestForGroup( + c.Request.Context(), + apiKey.Group, + forcedPlatform, + modelIDs, + ) + if err != nil { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest") + return + } + etag := service.CodexModelsManifestETag(body) + c.Header("ETag", etag) + if service.CodexModelsManifestETagMatches(c.GetHeader("If-None-Match"), etag) { + c.Status(http.StatusNotModified) + c.Writer.WriteHeaderNow() + return + } + c.Data(http.StatusOK, "application/json", body) +} + +func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *service.Group, platformOverride string) []string { + if h == nil || h.gatewayService == nil || group == nil { + return nil + } + + groupID := &group.ID + platform := strings.TrimSpace(platformOverride) + if platform == "" { + platform = group.Platform + } + if platform == service.PlatformComposite { + availableModels := h.compositeAvailableModels(ctx, groupID) + fallbackModels := defaultCodexModelIDsForPlatform(service.PlatformComposite) + if group.CustomModelsListEnabled() { + return filterModelsByCustomList(availableModels, fallbackModels, group.ModelsListConfig.Models) + } + if len(availableModels) > 0 { + return availableModels + } + return fallbackModels + } + + availableModels := h.gatewayService.GetAvailableModels(ctx, groupID, platform) + fallbackModels := defaultCodexModelIDsForPlatform(platform) + if group.CustomModelsListEnabled() { + return filterModelsByCustomList( + customModelsListSource(platform, availableModels, fallbackModels), + fallbackModels, + group.ModelsListConfig.Models, + ) + } + if len(availableModels) > 0 { + return availableModels + } + return fallbackModels +} + func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64) []string { if h == nil || h.gatewayService == nil { return nil @@ -1340,9 +1413,26 @@ func customModelsListAllowsModel(availablePatterns []string, model string) bool return true } } + normalizedClaudeModel := claude.NormalizeModelID(strings.TrimSuffix(model, "-thinking")) + if normalizedClaudeModel != model { + for _, pattern := range availablePatterns { + if pattern == normalizedClaudeModel { + return true + } + } + } return false } +func defaultCodexModelIDsForPlatform(platform string) []string { + switch platform { + case service.PlatformDeepseek: + return []string{"deepseek-v4-pro", "deepseek-v4-flash"} + default: + return defaultModelIDsForPlatform(platform) + } +} + func defaultModelIDsForPlatform(platform string) []string { switch platform { case service.PlatformOpenAI: @@ -1361,14 +1451,7 @@ func defaultModelIDsForPlatform(platform string) []string { } return ids case service.PlatformAnthropic: - ids := make([]string, 0, len(claude.DefaultModels)+len(antigravity.DefaultModels())) - for _, model := range claude.DefaultModels { - ids = append(ids, model.ID) - } - for _, model := range antigravity.DefaultModels() { - ids = append(ids, model.ID) - } - return mergeModelIDs(ids, nil) + return claude.DefaultModelIDs() case service.PlatformGrok: return xai.DefaultModelIDs() case service.PlatformComposite: diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index 313b3216fb..05069bf25a 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -25,6 +25,22 @@ type gatewayModelsResponseForTest struct { Data []gatewayModelItemForTest `json:"data"` } +type codexModelsResponseForTest struct { + Models []struct { + Slug string `json:"slug"` + SupportedReasoningLevels []codexReasoningLevelForTest `json:"supported_reasoning_levels"` + InputModalities []string `json:"input_modalities"` + ModelMessages map[string]json.RawMessage `json:"model_messages"` + TruncationPolicy map[string]json.RawMessage `json:"truncation_policy"` + AvailabilityNUX json.RawMessage `json:"availability_nux"` + Upgrade json.RawMessage `json:"upgrade"` + } `json:"models"` +} + +type codexReasoningLevelForTest struct { + Effort string `json:"effort"` +} + type gatewayModelItemForTest struct { ID string `json:"id"` Object string `json:"object"` @@ -52,6 +68,10 @@ func (s *gatewayModelsAccountRepoStub) ListSchedulableByGroupID(ctx context.Cont return out, nil } +func (s *gatewayModelsAccountRepoStub) ListByGroup(ctx context.Context, groupID int64) ([]service.Account, error) { + return s.ListSchedulableByGroupID(ctx, groupID) +} + func newGatewayModelsHandlerForTest(repo service.AccountRepository) *GatewayHandler { return &GatewayHandler{ gatewayService: service.NewGatewayService( @@ -70,6 +90,247 @@ func TestDefaultModelIDsForCompositeIncludesAntigravityDefaults(t *testing.T) { require.Contains(t, compositeIDs, antigravityIDs[0]) } +// Scenario: Anthropic defaults contain only Claude while Antigravity keeps its own Gemini models. +func TestDefaultModelIDsForAnthropicExcludeAntigravityGemini(t *testing.T) { + anthropicIDs := defaultModelIDsForPlatform(service.PlatformAnthropic) + require.Contains(t, anthropicIDs, "claude-opus-4-6") + require.NotContains(t, anthropicIDs, "gemini-2.5-flash") + + antigravityIDs := defaultModelIDsForPlatform(service.PlatformAntigravity) + require.Contains(t, antigravityIDs, "gemini-2.5-flash") +} + +// Scenario: non-OpenAI groups return a Codex manifest instead of a standard model list. +func TestGatewayCodexModels_NonOpenAIGroupsUseMappedModels(t *testing.T) { + tests := []struct { + name string + platform string + model string + efforts []string + modalities []string + }{ + { + name: "Grok", + platform: service.PlatformGrok, + model: "grok-4.6", + efforts: []string{"low", "medium", "high", "xhigh"}, + modalities: []string{"text", "image"}, + }, + { + name: "DeepSeek", + platform: service.PlatformDeepseek, + model: "deepseek-v4-pro", + efforts: []string{"low", "high", "max"}, + modalities: []string{"text"}, + }, + { + name: "provider-qualified Claude", + platform: service.PlatformAnthropic, + model: "anthropic/claude-sonnet-4-6", + efforts: []string{"low", "medium", "high", "max"}, + modalities: []string{"text"}, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + groupID := int64(100 + index) + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: tt.platform, + Credentials: map[string]any{ + "model_mapping": map[string]any{tt.model: tt.model}, + }, + }, + }, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: tt.platform}, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Len(t, got.Models, 1) + require.Equal(t, tt.model, got.Models[0].Slug) + require.NotEmpty(t, got.Models[0].ModelMessages) + require.NotEmpty(t, got.Models[0].TruncationPolicy) + require.NotNil(t, got.Models[0].AvailabilityNUX) + require.NotNil(t, got.Models[0].Upgrade) + require.Equal(t, tt.efforts, codexReasoningEffortsForTest(got.Models[0].SupportedReasoningLevels)) + require.Equal(t, tt.modalities, got.Models[0].InputModalities) + }) + } +} + +// Scenario: Composite manifests aggregate only administrator-configured models. +func TestGatewayCodexModels_CompositeUsesCompleteEffectiveModelList(t *testing.T) { + gin.SetMode(gin.TestMode) + const groupID int64 = 120 + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 3, + Platform: service.PlatformOpenAI, + Status: service.StatusActive, + Schedulable: true, + Credentials: map[string]any{}, + }, + { + ID: 1, + Platform: service.PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-5.5": "gpt-5.5"}, + }, + }, + { + ID: 2, + Platform: service.PlatformGrok, + Credentials: map[string]any{ + "model_mapping": map[string]any{"grok-4.6": "grok-4.6"}, + }, + }, + }, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"gpt-5.5", "grok-4.6"}, codexModelSlugsForTest(got.Models)) +} + +func TestGatewayCodexModels_GeneratedManifestUsesFinalBodyETag(t *testing.T) { + gin.SetMode(gin.TestMode) + const groupID int64 = 122 + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: {{ + ID: 1, + Platform: service.PlatformDeepseek, + Credentials: map[string]any{ + "model_mapping": map[string]any{"deepseek-v4-pro": "deepseek-v4-pro"}, + }, + }}, + }, + }) + group := &service.Group{ID: groupID, Platform: service.PlatformDeepseek} + + first := httptest.NewRecorder() + firstContext, _ := gin.CreateTestContext(first) + firstContext.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + firstContext.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{Group: group}) + h.CodexModels(firstContext) + + require.Equal(t, http.StatusOK, first.Code) + etag := first.Header().Get("ETag") + require.NotEmpty(t, etag) + require.Equal(t, service.CodexModelsManifestETag(first.Body.Bytes()), etag) + + second := httptest.NewRecorder() + secondContext, _ := gin.CreateTestContext(second) + secondContext.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + secondContext.Request.Header.Set("If-None-Match", "W/"+etag) + secondContext.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{Group: group}) + h.CodexModels(secondContext) + + require.Equal(t, http.StatusNotModified, second.Code) + require.Empty(t, second.Body.Bytes()) + require.Equal(t, etag, second.Header().Get("ETag")) +} + +// Scenario: group models_list_config limits the generated Codex manifest. +func TestGatewayCodexModels_CustomModelsListFiltersCompositeManifest(t *testing.T) { + gin.SetMode(gin.TestMode) + const groupID int64 = 121 + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-5.5": "gpt-5.5"}, + }, + }, + { + ID: 2, + Platform: service.PlatformGrok, + Credentials: map[string]any{ + "model_mapping": map[string]any{"grok-4.6": "grok-4.6"}, + }, + }, + }, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformComposite, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"grok-4.6"}, + }, + }, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"grok-4.6"}, codexModelSlugsForTest(got.Models)) +} + +func codexModelSlugsForTest(models []struct { + Slug string `json:"slug"` + SupportedReasoningLevels []codexReasoningLevelForTest `json:"supported_reasoning_levels"` + InputModalities []string `json:"input_modalities"` + ModelMessages map[string]json.RawMessage `json:"model_messages"` + TruncationPolicy map[string]json.RawMessage `json:"truncation_policy"` + AvailabilityNUX json.RawMessage `json:"availability_nux"` + Upgrade json.RawMessage `json:"upgrade"` +}) []string { + slugs := make([]string, 0, len(models)) + for _, model := range models { + slugs = append(slugs, model.Slug) + } + return slugs +} + +func codexReasoningEffortsForTest(levels []codexReasoningLevelForTest) []string { + efforts := make([]string, 0, len(levels)) + for _, level := range levels { + efforts = append(efforts, level.Effort) + } + return efforts +} + func TestGatewayModels_GeminiGroupFallsBackToGeminiModels(t *testing.T) { gin.SetMode(gin.TestMode) @@ -193,6 +454,62 @@ func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) { require.Equal(t, []string{"gemini-2.5-flash"}, modelIDsForTest(got.Data)) } +// Scenario: a Composite group with only Anthropic accounts must not inherit Antigravity Gemini defaults. +func TestGatewayCodexModels_CompositeAnthropicDoesNotAdvertiseAntigravityDefaults(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(64) + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: {{ID: 1, Platform: service.PlatformAnthropic}}, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + slugs := codexModelSlugsForTest(got.Models) + require.Contains(t, slugs, "claude-opus-4-6") + require.NotContains(t, slugs, "gemini-2.5-flash") +} + +// Scenario: Antigravity retains its own Claude and Gemini defaults inside Composite groups. +func TestGatewayModels_CompositeAntigravityAdvertisesAntigravityDefaults(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(65) + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: {{ID: 1, Platform: service.PlatformAntigravity}}, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + ids := modelIDsForTest(got.Data) + require.Contains(t, ids, "claude-opus-4-6") + require.Contains(t, ids, "gemini-2.5-flash") +} + func TestGatewayModels_CustomModelsListDisabledKeepsOriginalModels(t *testing.T) { gin.SetMode(gin.TestMode) @@ -457,6 +774,89 @@ func TestDefaultModelIDsForPlatform_CNProvidersKeepClaudeDefaults(t *testing.T) } } +func TestDefaultCodexModelIDsForPlatform_DeepSeekUsesDeepSeekModels(t *testing.T) { + require.Equal(t, []string{"deepseek-v4-pro", "deepseek-v4-flash"}, defaultCodexModelIDsForPlatform(service.PlatformDeepseek)) + require.Equal(t, defaultModelIDsForPlatform(service.PlatformAnthropic), defaultCodexModelIDsForPlatform(service.PlatformAnthropic)) +} + +func TestGatewayCodexModels_DeepSeekWithoutMappingUsesDeepSeekDefaults(t *testing.T) { + gin.SetMode(gin.TestMode) + const groupID int64 = 130 + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformDeepseek, + Status: service.StatusActive, + Schedulable: true, + Credentials: map[string]any{}, + }, + }, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.150.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformDeepseek}, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + slugs := make([]string, 0, len(got.Models)) + for _, model := range got.Models { + slugs = append(slugs, model.Slug) + } + require.Contains(t, slugs, "deepseek-v4-pro") + require.Contains(t, slugs, "deepseek-v4-flash") + require.NotContains(t, slugs, "claude-sonnet-4-6") + require.NotContains(t, slugs, "claude-opus-4-6") +} + +func TestGatewayCodexModels_OmitsWildcardMappingKeys(t *testing.T) { + gin.SetMode(gin.TestMode) + const groupID int64 = 131 + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformDeepseek, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "foo-*": "deepseek-v4-pro", + "deepseek-v4-pro": "deepseek-v4-pro", + }, + }, + }, + }, + }, + }) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.150.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformDeepseek}, + }) + + h.CodexModels(c) + + require.Equal(t, http.StatusOK, rec.Code) + var got codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + slugs := make([]string, 0, len(got.Models)) + for _, model := range got.Models { + slugs = append(slugs, model.Slug) + } + require.Equal(t, []string{"deepseek-v4-pro"}, slugs) +} + func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMapping(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index 2714dc294f..3a38ce5154 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -15,9 +15,9 @@ 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. ChatGPT manifests are proxied verbatim; custom API key manifests receive -// provider-compatibility normalization and use a short-lived, asynchronously -// revalidated cache to tolerate canceled client requests. +// here. Groups with explicit account model mappings are generated locally; +// otherwise ChatGPT manifests are proxied verbatim and custom API key manifests +// receive provider-compatibility normalization plus short-lived caching. func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { if c.Request.Context().Err() != nil { return @@ -32,6 +32,24 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { return } + ifNoneMatch := c.GetHeader("If-None-Match") + configuredManifest, configured, err := h.gatewayService.BuildGroupConfiguredCodexModelsManifest( + c.Request.Context(), + apiKey.Group, + ifNoneMatch, + ) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest") + return + } + if configured { + writeCodexModelsManifestResponse(c, configuredManifest) + return + } + maxAccountSwitches := h.maxAccountSwitches if maxAccountSwitches <= 0 { maxAccountSwitches = 3 @@ -56,7 +74,9 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { // 让 ops 错误日志携带实际选中的上游账号,便于定位失效账号(#4544)。 setOpsSelectedAccount(c, account.ID, account.Platform) - manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) + // The client ETag represents the final group-specific body, so fetch the + // source manifest before applying local filtering and alias metadata. + manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), "") if err != nil { if c.Request.Context().Err() != nil { return @@ -70,18 +90,31 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) return } + if err := h.gatewayService.CompleteAPIKeyCodexModelsManifestForClient(manifest, account); err != nil { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to complete Codex models manifest") + return + } + if err := h.gatewayService.MergeGroupConfiguredCodexModels(c.Request.Context(), apiKey.Group, manifest, ifNoneMatch); err != nil { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest") + 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) + writeCodexModelsManifestResponse(c, manifest) return } } + +func writeCodexModelsManifestResponse(c *gin.Context, manifest *service.CodexModelsManifest) { + if manifest.ETag != "" { + c.Header("ETag", manifest.ETag) + } + if manifest.NotModified { + c.Status(http.StatusNotModified) + c.Writer.WriteHeaderNow() + 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 index f8d1fada6c..9eabd3e5b2 100644 --- a/backend/internal/handler/openai_codex_models_handler_test.go +++ b/backend/internal/handler/openai_codex_models_handler_test.go @@ -2,6 +2,7 @@ package handler import ( "context" + "encoding/json" "errors" "fmt" "io" @@ -16,6 +17,7 @@ import ( middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" ) type codexModelsFailoverAccountRepo struct { @@ -43,6 +45,14 @@ func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Cont return accounts, nil } +func (r codexModelsFailoverAccountRepo) ListSchedulableByGroupID(_ context.Context, _ int64) ([]service.Account, error) { + return append([]service.Account(nil), r.accounts...), nil +} + +func (r codexModelsFailoverAccountRepo) ListByGroup(_ context.Context, _ int64) ([]service.Account, error) { + return append([]service.Account(nil), r.accounts...), nil +} + type codexModelsFailoverHTTPUpstream struct { service.HTTPUpstream mu sync.Mutex @@ -116,6 +126,236 @@ func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { } } +func TestCodexModelsAppliesLocalFiltersBeforeClientETag(t *testing.T) { + gin.SetMode(gin.TestMode) + groupID := int64(43) + repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{ + { + ID: 1, + Name: "custom-openai", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://upstream.example/v1", + }, + }, + }} + upstream := &codexModelsFailoverHTTPUpstream{ + firstBody: `{"object":"list","data":[{"id":"codex-auto-review"},{"id":"gpt-5.6"}]}`, + } + gatewayService := service.NewOpenAIGatewayService( + repo, + nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + handler := &OpenAIGatewayHandler{gatewayService: gatewayService} + group := &service.Group{ + ID: groupID, + Platform: service.PlatformOpenAI, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"codex-auto-review", "gpt-5.6"}, + }, + } + + first := performCodexModelsRequestForGroup(t, handler, group, "") + if first.Code != http.StatusOK { + t.Fatalf("first status: got %d, want %d; body=%s", first.Code, http.StatusOK, first.Body.String()) + } + if body := first.Body.String(); !strings.Contains(body, "codex-auto-review") || !strings.Contains(body, "gpt-5.6") { + t.Fatalf("first body did not include the explicitly selected models: %s", body) + } + oldETag := first.Header().Get("ETag") + if oldETag == "" { + t.Fatal("first response did not include an ETag") + } + + group.ModelsListConfig.Enabled = false + second := performCodexModelsRequestForGroup(t, handler, group, oldETag) + if second.Code != http.StatusOK { + t.Fatalf("second status: got %d, want %d; body=%s", second.Code, http.StatusOK, second.Body.String()) + } + if body := second.Body.String(); strings.Contains(body, "codex-auto-review") || !strings.Contains(body, "gpt-5.6") { + t.Fatalf("second body was not the filtered manifest: %s", body) + } + if newETag := second.Header().Get("ETag"); newETag == "" || newETag == oldETag { + t.Fatalf("second ETag: got %q, want a new final-body ETag", newETag) + } + + third := performCodexModelsRequestForGroup(t, handler, group, second.Header().Get("ETag")) + if third.Code != http.StatusNotModified { + t.Fatalf("third status: got %d, want %d; body=%s", third.Code, http.StatusNotModified, third.Body.String()) + } + if third.Body.Len() != 0 { + t.Fatalf("third body: got %q, want empty", third.Body.String()) + } +} + +func TestCodexModelsAPIKeyCacheDoesNotLeakGroupFilters(t *testing.T) { + gin.SetMode(gin.TestMode) + repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{ + { + ID: 1, + Name: "shared-api-key", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-shared", + "base_url": "https://upstream.example/v1", + }, + }, + }} + upstream := &codexModelsFailoverHTTPUpstream{ + firstBody: `{"object":"list","data":[{"id":"model-a"},{"id":"model-b"}]}`, + } + gatewayService := service.NewOpenAIGatewayService( + repo, + nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + handler := &OpenAIGatewayHandler{gatewayService: gatewayService} + groupA := &service.Group{ + ID: 91, + Platform: service.PlatformOpenAI, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"model-a"}, + }, + } + groupB := &service.Group{ + ID: 92, + Platform: service.PlatformOpenAI, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"model-b"}, + }, + } + + firstA := performCodexModelsRequestForGroup(t, handler, groupA, "") + require.Equal(t, http.StatusOK, firstA.Code, firstA.Body.String()) + require.Equal(t, []string{"model-a"}, codexHandlerManifestSlugs(t, firstA)) + + firstB := performCodexModelsRequestForGroup(t, handler, groupB, "") + require.Equal(t, http.StatusOK, firstB.Code, firstB.Body.String()) + require.Equal(t, []string{"model-b"}, codexHandlerManifestSlugs(t, firstB)) + + etagA := firstA.Header().Get("ETag") + require.NotEmpty(t, etagA) + + var wg sync.WaitGroup + results := make([]*httptest.ResponseRecorder, 8) + for i := range results { + wg.Add(1) + go func(index int) { + defer wg.Done() + if index%2 == 0 { + results[index] = performCodexModelsRequestForGroup(t, handler, groupA, etagA) + return + } + results[index] = performCodexModelsRequestForGroup(t, handler, groupB, "") + }(i) + } + wg.Wait() + + sawGroupB := false + for _, recorder := range results { + require.NotNil(t, recorder) + switch recorder.Code { + case http.StatusNotModified: + require.Empty(t, recorder.Body.Bytes()) + case http.StatusOK: + slugs := codexHandlerManifestSlugs(t, recorder) + if len(slugs) == 1 && slugs[0] == "model-b" { + sawGroupB = true + continue + } + require.Equal(t, []string{"model-a"}, slugs) + default: + t.Fatalf("unexpected status %d body=%s", recorder.Code, recorder.Body.String()) + } + } + require.True(t, sawGroupB) +} + +// Scenario: OpenAI 分组内混用 OAuth 和第三方 API Key 时,管理员模型配置优先。 +func TestCodexModelsUsesConfiguredModelsBeforeUpstreamDiscovery(t *testing.T) { + gin.SetMode(gin.TestMode) + groupID := int64(44) + repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{ + { + ID: 1, + Name: "ark-compatible", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Priority: 0, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-ark", + "base_url": "https://ark.example/v1", + "model_mapping": map[string]any{ + "glm-5.3": "glm-5.3", + }, + }, + }, + { + ID: 2, + Name: "chatgpt-oauth", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + Schedulable: true, + Priority: 1, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-test", + }, + }, + }} + upstream := &codexModelsFailoverHTTPUpstream{firstStatus: http.StatusNotFound} + gatewayService := service.NewOpenAIGatewayService( + repo, + nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + handler := &OpenAIGatewayHandler{gatewayService: gatewayService} + + recorder := performCodexModelsRequestForGroup(t, handler, &service.Group{ + ID: groupID, + Platform: service.PlatformOpenAI, + }, "") + + if got := upstream.calls(); len(got) != 0 { + t.Fatalf("upstream account calls: got %v, want none", got) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + var envelope struct { + Models []map[string]any `json:"models"` + } + if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String()) + } + if len(envelope.Models) != 1 || envelope.Models[0]["slug"] != "glm-5.3" { + t.Fatalf("models: got %v, want only glm-5.3", envelope.Models) + } + if _, ok := envelope.Models[0]["supported_reasoning_levels"]; !ok { + t.Fatalf("configured model is missing the Codex descriptor contract: %v", envelope.Models[0]) + } +} + func TestCompositeCodexModelsReusesExistingManifestSelection(t *testing.T) { handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) @@ -148,8 +388,23 @@ func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) { 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) + requireCompleteCodexModelsHandlerResponse(t, recorder, "gpt-5.6-sol") + }) + } +} + +// Scenario: an API-key upstream without /models is excluded only for this discovery request. +func TestCodexModelsFailsOverWhenAPIKeyModelsEndpointIsUnavailable(t *testing.T) { + for _, status := range []int{http.StatusNotFound, http.StatusMethodNotAllowed} { + 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()) } }) } @@ -183,9 +438,7 @@ func TestCodexModelsFailsOverFromInvalidManifestEnvelope(t *testing.T) { 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) - } + requireCompleteCodexModelsHandlerResponse(t, recorder, "gpt-5.6-sol") } func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) { @@ -193,7 +446,6 @@ func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) { http.StatusBadRequest, http.StatusUnauthorized, http.StatusForbidden, - http.StatusNotFound, 600, } for _, status := range statuses { @@ -300,23 +552,79 @@ func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount } func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder { - return performCodexModelsRequestForPlatform(t, handler, groupID, service.PlatformOpenAI) + return performCodexModelsRequestForGroup(t, handler, &service.Group{ID: groupID, Platform: service.PlatformOpenAI}, "") } func performCodexModelsRequestForPlatform(t *testing.T, handler *OpenAIGatewayHandler, groupID int64, platform string) *httptest.ResponseRecorder { + return performCodexModelsRequestForGroup(t, handler, &service.Group{ID: groupID, Platform: platform}, "") +} + +func performCodexModelsRequestForGroup(t *testing.T, handler *OpenAIGatewayHandler, group *service.Group, etag string) *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) + if etag != "" { + c.Request.Header.Set("If-None-Match", etag) + } c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ - GroupID: &groupID, - Group: &service.Group{ID: groupID, Platform: platform}, + GroupID: &group.ID, + Group: group, }) handler.CodexModels(c) return recorder } +func codexHandlerManifestSlugs(t *testing.T, recorder *httptest.ResponseRecorder) []string { + t.Helper() + + var envelope struct { + Models []struct { + Slug string `json:"slug"` + } `json:"models"` + } + if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String()) + } + slugs := make([]string, 0, len(envelope.Models)) + for _, model := range envelope.Models { + slugs = append(slugs, model.Slug) + } + return slugs +} + +func requireCompleteCodexModelsHandlerResponse(t *testing.T, recorder *httptest.ResponseRecorder, slug string) { + t.Helper() + + var envelope struct { + Models []map[string]any `json:"models"` + } + if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String()) + } + if len(envelope.Models) != 1 { + t.Fatalf("models count: got %d, want 1; body=%s", len(envelope.Models), recorder.Body.String()) + } + model := envelope.Models[0] + if got := model["slug"]; got != slug { + t.Fatalf("slug: got %v, want %q", got, slug) + } + if levels, ok := model["supported_reasoning_levels"].([]any); !ok || len(levels) == 0 { + t.Fatalf("supported_reasoning_levels must be populated: %v", model["supported_reasoning_levels"]) + } + if messages, ok := model["model_messages"].(map[string]any); !ok || messages["instructions_template"] == "" { + t.Fatalf("model_messages.instructions_template must be populated: %v", model["model_messages"]) + } + if policy, ok := model["truncation_policy"].(map[string]any); !ok || len(policy) == 0 { + t.Fatalf("truncation_policy must be populated: %v", model["truncation_policy"]) + } + modalities, ok := model["input_modalities"].([]any) + if !ok || len(modalities) != 1 || modalities[0] != "text" { + t.Fatalf("custom OpenAI-compatible endpoint modalities: got %v, want [text]", model["input_modalities"]) + } +} + func equalInt64Slices(got, want []int64) bool { if len(got) != len(want) { return false diff --git a/backend/internal/pkg/apicompat/responses_client_tools.go b/backend/internal/pkg/apicompat/responses_client_tools.go index b1e2c412f1..5e831c1c52 100644 --- a/backend/internal/pkg/apicompat/responses_client_tools.go +++ b/backend/internal/pkg/apicompat/responses_client_tools.go @@ -235,12 +235,12 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) (b typed["type"] = "function_call" typed["arguments"] = customToolCallArguments(stringValue(typed["input"])) delete(typed, "input") - dropInvalidLoweredFunctionItemID(typed) + normalizeLoweredFunctionItemID(typed) changed = true } case "custom_tool_call_output": typed["type"] = "function_call_output" - dropInvalidLoweredFunctionItemID(typed) + normalizeLoweredFunctionItemID(typed) normalizeClientToolOutput(typed) changed = true case "tool_search_call": @@ -249,7 +249,7 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) (b typed["name"] = toolSearchProxyName typed["arguments"] = rawObjectString(typed["arguments"]) delete(typed, "execution") - dropInvalidLoweredFunctionItemID(typed) + normalizeLoweredFunctionItemID(typed) changed = true } case "tool_search_output": @@ -259,7 +259,7 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) (b return fmt.Errorf("tool_search_output requires a non-empty string call_id before it can be lowered to function_call_output") } typed["type"] = "function_call_output" - dropInvalidLoweredFunctionItemID(typed) + normalizeLoweredFunctionItemID(typed) if err := normalizeToolSearchOutput(typed); err != nil { return err } @@ -280,14 +280,73 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) (b return changed, nil } -// dropInvalidLoweredFunctionItemID removes Codex client-only item IDs such as +// normalizeLoweredFunctionItemID reconciles Codex client-only item IDs such as // ctc_*, ctco_*, tsc_*, and tso_* after their item type is lowered to the -// function protocol. Function upstreams validate these IDs with the fc prefix; -// call_id, which is preserved separately, is the tool call/output pairing key. -func dropInvalidLoweredFunctionItemID(item map[string]any) { +// function protocol, which validates these IDs with the fc prefix. A ctc_/tsc_ +// call ID maps back to the fc_ ID it was raised from; the remaining IDs have no +// function-protocol counterpart and are dropped. call_id, which is preserved +// separately, stays the tool call/output pairing key either way. +func normalizeLoweredFunctionItemID(item map[string]any) { id := strings.TrimSpace(stringValue(item["id"])) - if id != "" && !strings.HasPrefix(id, "fc") { - delete(item, "id") + if id == "" || strings.HasPrefix(id, "fc") { + return + } + // A ctc_/tsc_ call ID on a lowered call item is the one we minted from the + // upstream's own fc_ ID when the item was raised, so map it back instead of + // dropping it. Anything else has no known function-protocol counterpart. + if recovered := retypedResponsesToolCallItemID(id, "function_call"); recovered != id { + item["id"] = recovered + return + } + delete(item, "id") +} + +// responsesToolCallItemIDPrefixes lists the Responses item ID prefixes that are +// tied to a specific tool-call item type. +var responsesToolCallItemIDPrefixes = []string{"fc_", "ctc_", "tsc_"} + +// responsesToolCallItemIDPrefix reports the ID prefix the Responses API +// validates for itemType, or "" when the type constrains no prefix. +func responsesToolCallItemIDPrefix(itemType string) string { + switch itemType { + case "custom_tool_call": + return "ctc_" + case "tool_search_call": + return "tsc_" + case "function_call": + return "fc_" + default: + return "" + } +} + +// retypedResponsesToolCallItemID re-prefixes an upstream item ID so it agrees +// with the item type we raise it back to. A function-only upstream answers a +// lowered custom tool with an fc_ item ID; emitting that ID on the restored +// custom_tool_call poisons the client's history, because a later replay of the +// same item to an upstream that validates IDs fails with +// "Invalid 'input[N].id' ... Expected an ID that begins with 'ctc'". +// The suffix is preserved so the ID stays stable and unique per upstream item. +// IDs that carry no known tool-call prefix are left alone rather than guessed at. +func retypedResponsesToolCallItemID(id, itemType string) string { + want := responsesToolCallItemIDPrefix(itemType) + if want == "" || id == "" || strings.HasPrefix(id, want) { + return id + } + for _, known := range responsesToolCallItemIDPrefixes { + if known != want && strings.HasPrefix(id, known) { + return want + strings.TrimPrefix(id, known) + } + } + return id +} + +// retypeResponsesToolCallItemID applies retypedResponsesToolCallItemID to a +// decoded item map. +func retypeResponsesToolCallItemID(item map[string]any, itemType string) { + id := strings.TrimSpace(stringValue(item["id"])) + if retyped := retypedResponsesToolCallItemID(id, itemType); retyped != id { + item["id"] = retyped } } @@ -432,12 +491,14 @@ func restoreClientToolValue(value any, adapter *ResponsesClientToolMapping) bool name := strings.TrimSpace(stringValue(typed["name"])) if adapter.CustomTools[name] { typed["type"] = "custom_tool_call" + retypeResponsesToolCallItemID(typed, "custom_tool_call") typed["input"] = extractCustomToolCallInput(rawObjectString(typed["arguments"])) delete(typed, "arguments") delete(typed, "namespace") changed = true } else if adapter.ToolSearch && name == toolSearchProxyName { typed["type"] = "tool_search_call" + retypeResponsesToolCallItemID(typed, "tool_search_call") typed["execution"] = "client" typed["arguments"] = json.RawMessage(toolSearchCallArgumentsJSON(rawObjectString(typed["arguments"]))) delete(typed, "name") @@ -464,12 +525,15 @@ type ResponsesClientToolStreamRestorer struct { } type responsesClientToolStreamCall struct { - kind string - name string - callID string - itemID string - outputIdx int - arguments strings.Builder + kind string + name string + // callID and itemID stay as the upstream sent them so later upstream + // events keep matching this call; clientItemID is what we emit. + callID string + itemID string + clientItemID string + outputIdx int + arguments strings.Builder } func NewResponsesClientToolStreamRestorer(mapping ResponsesClientToolMapping) *ResponsesClientToolStreamRestorer { @@ -508,6 +572,9 @@ func (r *ResponsesClientToolStreamRestorer) Restore(event ResponsesStreamEvent) event.Item.Arguments = "{}" event.Item.Namespace = "" } + if call.clientItemID != "" { + event.Item.ID = call.clientItemID + } } emit(r.restoreNamespaceEvent(event)) case "response.function_call_arguments.delta": @@ -525,9 +592,9 @@ func (r *ResponsesClientToolStreamRestorer) Restore(event ResponsesStreamEvent) if call.kind == "custom" { input := extractCustomToolCallInput(call.arguments.String()) if input != "" { - emit(ResponsesStreamEvent{Type: "response.custom_tool_call_input.delta", OutputIndex: call.outputIdx, ItemID: call.itemID, Delta: input}) + emit(ResponsesStreamEvent{Type: "response.custom_tool_call_input.delta", OutputIndex: call.outputIdx, ItemID: call.clientItemID, Delta: input}) } - emit(ResponsesStreamEvent{Type: "response.custom_tool_call_input.done", OutputIndex: call.outputIdx, ItemID: call.itemID, CallID: call.callID, Name: call.name, Input: input}) + emit(ResponsesStreamEvent{Type: "response.custom_tool_call_input.done", OutputIndex: call.outputIdx, ItemID: call.clientItemID, CallID: call.callID, Name: call.name, Input: input}) } return out } @@ -548,6 +615,9 @@ func (r *ResponsesClientToolStreamRestorer) Restore(event ResponsesStreamEvent) } event.Item.Namespace = "" } + if call.clientItemID != "" { + event.Item.ID = call.clientItemID + } delete(r.calls, call.itemID) delete(r.calls, call.callID) delete(r.byOutput, call.outputIdx) @@ -684,6 +754,15 @@ func (r *ResponsesClientToolStreamRestorer) resequenceRaw(payload []byte, sequen return [][]byte{encoded}, true, nil } +// responsesClientToolItemType maps a restorer call kind to the item type the +// client sees. +func responsesClientToolItemType(kind string) string { + if kind == "custom" { + return "custom_tool_call" + } + return "tool_search_call" +} + func (r *ResponsesClientToolStreamRestorer) recordItem(event ResponsesStreamEvent) *responsesClientToolStreamCall { if event.Item == nil || event.Item.Type != "function_call" { return nil @@ -704,7 +783,14 @@ func (r *ResponsesClientToolStreamRestorer) recordItem(event ResponsesStreamEven } call := r.calls[key] if call == nil { - call = &responsesClientToolStreamCall{kind: kind, name: name, callID: event.Item.CallID, itemID: event.Item.ID, outputIdx: event.OutputIndex} + call = &responsesClientToolStreamCall{ + kind: kind, + name: name, + callID: event.Item.CallID, + itemID: event.Item.ID, + clientItemID: retypedResponsesToolCallItemID(event.Item.ID, responsesClientToolItemType(kind)), + outputIdx: event.OutputIndex, + } r.calls[key] = call if call.callID != "" { r.calls[call.callID] = call @@ -758,11 +844,13 @@ func restoreResponsesOutputClientTools(outputs []ResponsesOutput, adapter *Respo } if adapter.CustomTools[output.Name] { output.Type = "custom_tool_call" + output.ID = retypedResponsesToolCallItemID(output.ID, output.Type) output.Input = extractCustomToolCallInput(output.Arguments) output.Arguments = "" output.Namespace = "" } else if adapter.ToolSearch && output.Name == toolSearchProxyName { output.Type = "tool_search_call" + output.ID = retypedResponsesToolCallItemID(output.ID, output.Type) output.Name = "" output.Namespace = "" } diff --git a/backend/internal/pkg/apicompat/responses_client_tools_item_id_helper_test.go b/backend/internal/pkg/apicompat/responses_client_tools_item_id_helper_test.go new file mode 100644 index 0000000000..66448ed661 --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_client_tools_item_id_helper_test.go @@ -0,0 +1,29 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRetypedResponsesToolCallItemID(t *testing.T) { + for _, tc := range []struct { + name string + id string + itemType string + want string + }{ + {"function id raised to custom", "fc_abc", "custom_tool_call", "ctc_abc"}, + {"function id raised to tool search", "fc_abc", "tool_search_call", "tsc_abc"}, + {"already correct is untouched", "ctc_abc", "custom_tool_call", "ctc_abc"}, + {"custom id lowered to function", "ctc_abc", "function_call", "fc_abc"}, + {"unknown prefix is left alone", "item_abc", "custom_tool_call", "item_abc"}, + {"unprefixed id is left alone", "abc", "custom_tool_call", "abc"}, + {"empty id stays empty", "", "custom_tool_call", ""}, + {"unconstrained item type is left alone", "fc_abc", "message", "fc_abc"}, + } { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.want, retypedResponsesToolCallItemID(tc.id, tc.itemType)) + }) + } +} diff --git a/backend/internal/pkg/apicompat/responses_client_tools_item_id_test.go b/backend/internal/pkg/apicompat/responses_client_tools_item_id_test.go new file mode 100644 index 0000000000..724f77b75d --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_client_tools_item_id_test.go @@ -0,0 +1,132 @@ +package apicompat + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +// A function-only upstream answers a lowered custom tool with an fc_ item ID. +// Emitting that ID on the restored custom_tool_call breaks the client the next +// time the item is replayed to an upstream that validates item IDs: +// "Invalid 'input[N].id': 'fc_...'. Expected an ID that begins with 'ctc'". +func TestRestoreResponsesClientToolPayload_RetypesToolCallItemIDs(t *testing.T) { + mapping := ResponsesClientToolMapping{ + CustomTools: map[string]bool{"exec": true}, ToolSearch: true, + NamespaceTools: map[string]ResponsesNamespaceName{"team__send": {Namespace: "team", Name: "send"}}, + } + payload := []byte(`{"id":"resp","output":[` + + `{"type":"function_call","id":"fc_abc123","call_id":"call_1","name":"exec","arguments":"{\"input\":\"dir\"}"},` + + `{"type":"function_call","id":"fc_def456","call_id":"call_2","name":"tool_search","arguments":"{\"query\":\"git\"}"},` + + `{"type":"function_call","id":"fc_ghi789","call_id":"call_3","name":"team__send","arguments":"{}"}]}`) + + restored, changed, err := RestoreResponsesClientToolPayload(payload, mapping) + require.NoError(t, err) + require.True(t, changed) + require.JSONEq(t, `{"id":"resp","output":[`+ + `{"type":"custom_tool_call","id":"ctc_abc123","call_id":"call_1","name":"exec","input":"dir"},`+ + `{"type":"tool_search_call","id":"tsc_def456","call_id":"call_2","execution":"client","arguments":{"query":"git"}},`+ + `{"type":"function_call","id":"fc_ghi789","call_id":"call_3","name":"send","namespace":"team","arguments":"{}"}]}`, + string(restored)) +} + +func TestRestoreResponsesOutputClientTools_RetypesToolCallItemIDs(t *testing.T) { + mapping := ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}, ToolSearch: true} + outputs := []ResponsesOutput{ + {Type: "function_call", ID: "fc_abc123", CallID: "call_1", Name: "exec", Arguments: `{"input":"dir"}`}, + {Type: "function_call", ID: "fc_def456", CallID: "call_2", Name: toolSearchProxyName, Arguments: `{"query":"git"}`}, + } + + restoreResponsesOutputClientTools(outputs, &mapping) + + require.Equal(t, "custom_tool_call", outputs[0].Type) + require.Equal(t, "ctc_abc123", outputs[0].ID) + require.Equal(t, "call_1", outputs[0].CallID) + require.Equal(t, "tool_search_call", outputs[1].Type) + require.Equal(t, "tsc_def456", outputs[1].ID) + require.Equal(t, "call_2", outputs[1].CallID) +} + +func TestResponsesClientToolStreamRestorer_RetypesCustomToolCallItemID(t *testing.T) { + const upstreamID = "fc_09f77ac43cf7db36016a8920e7934487" + const clientID = "ctc_09f77ac43cf7db36016a8920e7934487" + + restorer := NewResponsesClientToolStreamRestorer(ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}}) + + added := restorer.Restore(ResponsesStreamEvent{ + Type: "response.output_item.added", SequenceNumber: 0, OutputIndex: 0, + Item: &ResponsesOutput{Type: "function_call", ID: upstreamID, CallID: "call_1", Name: "exec", Status: "in_progress"}, + }) + require.Len(t, added, 1) + require.Equal(t, "custom_tool_call", added[0].Item.Type) + require.Equal(t, clientID, added[0].Item.ID) + + // Later upstream events still address the item by its upstream ID. + require.Empty(t, restorer.Restore(ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", SequenceNumber: 1, ItemID: upstreamID, Delta: `{"input":"di`, + })) + done := restorer.Restore(ResponsesStreamEvent{ + Type: "response.function_call_arguments.done", SequenceNumber: 2, ItemID: upstreamID, + CallID: "call_1", Name: "exec", Arguments: `{"input":"dir"}`, + }) + require.Len(t, done, 2) + require.Equal(t, "response.custom_tool_call_input.delta", done[0].Type) + require.Equal(t, clientID, done[0].ItemID) + require.Equal(t, "response.custom_tool_call_input.done", done[1].Type) + require.Equal(t, clientID, done[1].ItemID) + require.Equal(t, "call_1", done[1].CallID) + + closed := restorer.Restore(ResponsesStreamEvent{ + Type: "response.output_item.done", SequenceNumber: 3, OutputIndex: 0, + Item: &ResponsesOutput{Type: "function_call", ID: upstreamID, CallID: "call_1", Name: "exec", Arguments: `{"input":"dir"}`, Status: "completed"}, + }) + require.Len(t, closed, 1) + require.Equal(t, "custom_tool_call", closed[0].Item.Type) + require.Equal(t, clientID, closed[0].Item.ID) + require.Equal(t, "dir", closed[0].Item.Input) +} + +func TestResponsesClientToolStreamRestorer_RetypesToolSearchCallItemID(t *testing.T) { + restorer := NewResponsesClientToolStreamRestorer(ResponsesClientToolMapping{ToolSearch: true}) + + added := restorer.Restore(ResponsesStreamEvent{ + Type: "response.output_item.added", SequenceNumber: 0, OutputIndex: 0, + Item: &ResponsesOutput{Type: "function_call", ID: "fc_search1", CallID: "call_1", Name: toolSearchProxyName, Status: "in_progress"}, + }) + require.Len(t, added, 1) + require.Equal(t, "tool_search_call", added[0].Item.Type) + require.Equal(t, "tsc_search1", added[0].Item.ID) +} + +// The WS bridge replays restored items back to the upstream, so the ID we hand +// the client has to map back to the upstream's own fc_ ID on the way down. +func TestAdaptResponsesClientTools_RecoversRetypedToolCallItemID(t *testing.T) { + req := map[string]any{ + "tools": []any{ + map[string]any{"type": "custom", "name": "exec"}, + map[string]any{"type": "tool_search"}, + }, + "input": []any{ + map[string]any{"type": "custom_tool_call", "id": "ctc_upstream1", "call_id": "call_1", "name": "exec", "input": "dir"}, + map[string]any{"type": "tool_search_call", "id": "tsc_upstream2", "call_id": "call_2", "arguments": map[string]any{"query": "git"}}, + map[string]any{"type": "custom_tool_call_output", "id": "ctco_client", "call_id": "call_1", "output": "ok"}, + }, + } + + _, changed, err := AdaptResponsesClientTools(req) + require.NoError(t, err) + require.True(t, changed) + + input := requireResponsesClientToolValue[[]any](t, req["input"]) + require.Len(t, input, 3) + customCall := requireResponsesClientToolValue[map[string]any](t, input[0]) + require.Equal(t, "function_call", customCall["type"]) + require.Equal(t, "fc_upstream1", customCall["id"]) + searchCall := requireResponsesClientToolValue[map[string]any](t, input[1]) + require.Equal(t, "function_call", searchCall["type"]) + require.Equal(t, "fc_upstream2", searchCall["id"]) + // Output items have no function-protocol ID counterpart and stay dropped. + customOutput := requireResponsesClientToolValue[map[string]any](t, input[2]) + require.Equal(t, "function_call_output", customOutput["type"]) + require.NotContains(t, customOutput, "id") +} diff --git a/backend/internal/pkg/apicompat/responses_client_tools_test.go b/backend/internal/pkg/apicompat/responses_client_tools_test.go index 2b3460755f..bfe3fd2eb7 100644 --- a/backend/internal/pkg/apicompat/responses_client_tools_test.go +++ b/backend/internal/pkg/apicompat/responses_client_tools_test.go @@ -48,14 +48,14 @@ func TestAdaptResponsesClientTools_LowersDeclarationsHistoryChoiceAndNamespaces( input := requireResponsesClientToolValue[[]any](t, req["input"]) customCall := requireResponsesClientToolValue[map[string]any](t, input[0]) require.Equal(t, "function_call", customCall["type"]) - require.NotContains(t, customCall, "id") + require.Equal(t, "fc_client", customCall["id"]) require.JSONEq(t, `{"input":"dir"}`, requireResponsesClientToolValue[string](t, customCall["arguments"])) customOutput := requireResponsesClientToolValue[map[string]any](t, input[1]) require.Equal(t, "function_call_output", customOutput["type"]) require.NotContains(t, customOutput, "id") searchCall := requireResponsesClientToolValue[map[string]any](t, input[2]) require.Equal(t, "function_call", searchCall["type"]) - require.NotContains(t, searchCall, "id") + require.Equal(t, "fc_client", searchCall["id"]) require.Equal(t, toolSearchProxyName, searchCall["name"]) require.JSONEq(t, `{"query":"git"}`, requireResponsesClientToolValue[string](t, searchCall["arguments"])) searchOutput := requireResponsesClientToolValue[map[string]any](t, input[3]) diff --git a/backend/internal/pkg/claude/effort_catalog.go b/backend/internal/pkg/claude/effort_catalog.go new file mode 100644 index 0000000000..ab587a4c01 --- /dev/null +++ b/backend/internal/pkg/claude/effort_catalog.go @@ -0,0 +1,69 @@ +package claude + +import ( + "strings" + "unicode" +) + +var ( + effortLowMediumHigh = []string{"low", "medium", "high"} + effortLowMediumHighMax = []string{"low", "medium", "high", "max"} + effortLowMediumHighXHighMax = []string{"low", "medium", "high", "xhigh", "max"} +) + +var effortFamilies = []struct { + family string + levels []string +}{ + {family: "claude-mythos-preview", levels: effortLowMediumHighMax}, + {family: "claude-mythos-5", levels: effortLowMediumHighXHighMax}, + {family: "claude-fable-5", levels: effortLowMediumHighXHighMax}, + {family: "claude-sonnet-4-6", levels: effortLowMediumHighMax}, + {family: "claude-sonnet-5", levels: effortLowMediumHighXHighMax}, + {family: "claude-opus-4-8", levels: effortLowMediumHighXHighMax}, + {family: "claude-opus-4-7", levels: effortLowMediumHighXHighMax}, + {family: "claude-opus-4-6", levels: effortLowMediumHighMax}, + {family: "claude-opus-4-5", levels: effortLowMediumHigh}, + {family: "claude-opus-5", levels: effortLowMediumHighXHighMax}, +} + +// EffortLevelsForModel returns the output_config.effort values accepted by a +// Claude model, ordered from the lightest to the deepest reasoning level. +func EffortLevelsForModel(model string) []string { + id := normalizeEffortModelID(model) + for _, entry := range effortFamilies { + if id == entry.family || strings.HasPrefix(id, entry.family+"-") { + return append([]string(nil), entry.levels...) + } + } + return nil +} + +func normalizeEffortModelID(model string) string { + id := strings.ToLower(strings.TrimSpace(model)) + id = strings.TrimPrefix(id, "models/") + if slash := strings.IndexByte(id, '/'); slash >= 0 { + id = strings.TrimPrefix(strings.TrimSpace(id[slash+1:]), "models/") + } + id = strings.TrimPrefix(id, "anthropic.") + id = strings.TrimSuffix(id, "-thinking") + if mapped, ok := ModelIDReverseOverrides[id]; ok { + id = mapped + } + if len(id) >= 9 { + suffix := id[len(id)-9:] + if suffix[0] == '-' { + digits := true + for _, r := range suffix[1:] { + if !unicode.IsDigit(r) { + digits = false + break + } + } + if digits { + id = id[:len(id)-9] + } + } + } + return id +} diff --git a/backend/internal/pkg/claude/effort_catalog_test.go b/backend/internal/pkg/claude/effort_catalog_test.go new file mode 100644 index 0000000000..718e20eb57 --- /dev/null +++ b/backend/internal/pkg/claude/effort_catalog_test.go @@ -0,0 +1,29 @@ +package claude + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestEffortLevelsForModel(t *testing.T) { + t.Parallel() + + tests := []struct { + model string + want []string + }{ + {model: "claude-opus-4-6", want: []string{"low", "medium", "high", "max"}}, + {model: "anthropic/claude-sonnet-4-6", want: []string{"low", "medium", "high", "max"}}, + {model: "claude-opus-5", want: []string{"low", "medium", "high", "xhigh", "max"}}, + {model: "claude-opus-4-5-20251101", want: []string{"low", "medium", "high"}}, + {model: "claude-haiku-4-5-20251001", want: nil}, + {model: "gpt-5.6", want: nil}, + } + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.want, EffortLevelsForModel(tt.model)) + }) + } +} diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index 6fab34dc39..65634598f1 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -251,6 +251,30 @@ func IsGrokModelID(model string) bool { return false } +// IsGrokImagineModel reports whether model is a Grok Imagine image or video +// model. These media models cannot act as the primary Codex agent model. +func IsGrokImagineModel(model string) bool { + normalized := strings.ToLower(StripGrokProviderPrefix(model)) + if normalized == "" { + return false + } + if strings.HasPrefix(normalized, "imagine") { + return true + } + switch { + case normalized == "grok-imagine", + normalized == "grok-imagine-1", + normalized == "grok-imagine-edit", + normalized == "grok-video-1.5": + return true + case strings.HasPrefix(normalized, "grok-imagine-image"), + strings.HasPrefix(normalized, "grok-imagine-video"): + return true + default: + return false + } +} + // IsGrokTextResponsesModelID reports whether model is a known Grok text model // for the Responses API. Imagine image/video and unknown custom ids return false. func IsGrokTextResponsesModelID(model string) bool { diff --git a/backend/internal/pkg/xai/models_test.go b/backend/internal/pkg/xai/models_test.go index fb4860ee4d..f9d58e5a0c 100644 --- a/backend/internal/pkg/xai/models_test.go +++ b/backend/internal/pkg/xai/models_test.go @@ -58,6 +58,16 @@ func TestIsGrokModelID(t *testing.T) { require.False(t, IsGrokModelID("claude-sonnet-4")) } +func TestIsGrokImagineModel(t *testing.T) { + t.Parallel() + require.True(t, IsGrokImagineModel("grok-imagine-image")) + require.True(t, IsGrokImagineModel("grok-imagine-video-1.5-preview")) + require.True(t, IsGrokImagineModel("xai/grok-imagine-image-quality")) + require.True(t, IsGrokImagineModel("grok-video-1.5")) + require.False(t, IsGrokImagineModel("grok-4.6")) + require.False(t, IsGrokImagineModel("grok-build-0.1")) +} + func TestDefaultModelsIncludesGrok46(t *testing.T) { t.Parallel() ids := DefaultModelIDs() diff --git a/backend/internal/repository/channel_monitor_v2_aggregation.go b/backend/internal/repository/channel_monitor_v2_aggregation.go index 6ad2bbf9b3..9bec5877b5 100644 --- a/backend/internal/repository/channel_monitor_v2_aggregation.go +++ b/backend/internal/repository/channel_monitor_v2_aggregation.go @@ -255,7 +255,7 @@ WITH dedup AS ( -- group errors aggregate under platform 'composite', which is never an -- enabled config platform, and are filtered out of every monitor v2 query. lower(CASE - WHEN g.platform = 'composite' THEN COALESCE(NULLIF(TRIM(a.platform)), NULLIF(NULLIF(lower(TRIM(current_error.platform)), ''), 'composite'), 'unknown') + WHEN g.platform = 'composite' THEN COALESCE(NULLIF(TRIM(a.platform), ''), NULLIF(NULLIF(lower(TRIM(current_error.platform)), ''), 'composite'), 'unknown') ELSE COALESCE(NULLIF(TRIM(current_error.platform), ''), 'unknown') END) AS platform, COALESCE(current_error.group_id, 0) AS group_id, diff --git a/backend/internal/repository/channel_monitor_v2_repo_test.go b/backend/internal/repository/channel_monitor_v2_repo_test.go index 55c9a96460..82d5e61130 100644 --- a/backend/internal/repository/channel_monitor_v2_repo_test.go +++ b/backend/internal/repository/channel_monitor_v2_repo_test.go @@ -116,6 +116,8 @@ func TestChannelMonitorV2ErrorAggregationResolvesCompositePlatform(t *testing.T) require.Contains(t, query, "left join groups g on g.id = current_error.group_id") require.Contains(t, query, "left join accounts a on a.id = current_error.account_id") require.Contains(t, query, "a.platform") + require.Contains(t, query, "nullif(trim(a.platform), '')") + require.NotContains(t, query, "nullif(trim(a.platform))") } func TestChannelMonitorV2UsageSuccessExcludesCyberBillingRows(t *testing.T) { diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go index c072f09b27..3d02d6fbff 100644 --- a/backend/internal/repository/user_repo.go +++ b/backend/internal/repository/user_repo.go @@ -1145,12 +1145,17 @@ func (r *userRepository) ExistsByEmailAlias(ctx context.Context, email string) ( } func existsByEmailAliasWithClient(ctx context.Context, client *dbent.Client, email string) (bool, error) { + _, exists, err := emailAliasOwnerIDWithClient(ctx, client, email, 0) + return exists, err +} + +func emailAliasOwnerIDWithClient(ctx context.Context, client *dbent.Client, email string, currentUserID int64) (int64, bool, error) { if client == nil { - return false, nil + return 0, false, nil } probes := service.EmailAliasDedupProbes(email) if len(probes) == 0 { - return false, nil + return 0, false, nil } preds := make([]predicate.User, 0, 2*len(probes)) @@ -1164,20 +1169,82 @@ func existsByEmailAliasWithClient(ctx context.Context, client *dbent.Client, ema candidates, err := client.User.Query(). Where(dbuser.Or(preds...)). Limit(emailAliasCandidateLimit). - Select(dbuser.FieldEmail). - Strings(ctx) + Select(dbuser.FieldID, dbuser.FieldEmail). + All(ctx) if err != nil { - return false, err + return 0, false, err } // 探针会有过度匹配(点号只在 Gmail 家族无意义),最终判定必须回到完整归一化规则。 + // 返回“其他用户”优先于当前用户,避免历史重复数据让调用方误判为仅当前用户占用。 identity := service.NormalizeEmailForAliasDedup(email) + var selfID int64 + selfExists := false for _, candidate := range candidates { - if service.NormalizeEmailForAliasDedup(candidate) == identity { - return true, nil + if service.NormalizeEmailForAliasDedup(candidate.Email) != identity { + continue + } + if candidate.ID != 0 && candidate.ID != currentUserID { + return candidate.ID, true, nil + } + if candidate.ID == currentUserID { + selfID = candidate.ID + selfExists = true } } - return false, nil + return selfID, selfExists, nil +} + +// UpdateEmailWithAliasGuard 在调用方事务内更新主邮箱与密码哈希。 +// +// 邮箱换绑不能只依赖服务层前置查重:两个并发请求可能同时看到同一收件箱未被占用。 +// 这里先按“字面邮箱 + 收件箱身份”加锁,复查是否已被其他用户占用,再执行写入; +// PostgreSQL 使用事务级 advisory lock 跨实例互斥,测试内存库则由进程内锁兜底。 +func (r *userRepository) UpdateEmailWithAliasGuard( + ctx context.Context, + userID int64, + email string, + passwordHash string, +) error { + if userID <= 0 { + return service.ErrUserNotFound + } + if strings.TrimSpace(email) == "" || passwordHash == "" { + return fmt.Errorf("email identity update requires email and password hash") + } + tx := dbent.TxFromContext(ctx) + if tx == nil { + return fmt.Errorf("email identity update requires a transaction") + } + client := tx.Client() + + releaseEmailLock, err := lockRepositoryScopedKeys( + ctx, + client, + txAwareSQLExecutor(ctx, r.sql, r.client), + normalizedEmailUniquenessLockKey(email), + emailAliasUniquenessLockKey(email), + ) + if err != nil { + return err + } + defer releaseEmailLock() + + ownerID, exists, err := emailAliasOwnerIDWithClient(ctx, client, email, userID) + if err != nil { + return err + } + if exists && ownerID != userID { + return service.ErrEmailExists + } + + if _, err := client.User.UpdateOneID(userID). + SetEmail(email). + SetPasswordHash(passwordHash). + Save(ctx); err != nil { + return translatePersistenceError(err, service.ErrUserNotFound, service.ErrEmailExists) + } + return nil } // dotStrippedEmailExpr 渲染下面的表达式:去掉存量邮箱的大小写、首尾空白(与 diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 6c0bfd908d..fe753b2389 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -65,13 +65,13 @@ func RegisterGatewayRoutes( h.Gateway.CountTokens(c) } } + codexModelsHandler := func(c *gin.Context) { + dispatchCodexModelsGateway(c, h.OpenAIGateway.CodexModels, h.Gateway.CodexModels) + } modelsHandler := func(c *gin.Context) { if c.Query("client_version") != "" { - switch getGroupPlatform(c) { - case service.PlatformOpenAI, service.PlatformComposite: - h.OpenAIGateway.CodexModels(c) - return - } + codexModelsHandler(c) + return } h.Gateway.Models(c) } @@ -377,7 +377,7 @@ func RegisterGatewayRoutes( codexDirect.GET("/responses", func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) - codexDirect.GET("/models", h.OpenAIGateway.CodexModels) + codexDirect.GET("/models", codexModelsHandler) } // OpenAI Chat Completions API(不带v1前缀的别名)— auto-route based on group platform r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) { @@ -504,6 +504,14 @@ func RegisterGatewayRoutes( } +func dispatchCodexModelsGateway(c *gin.Context, openAIHandler, generatedHandler gin.HandlerFunc) { + if getGroupPlatform(c) == service.PlatformOpenAI { + openAIHandler(c) + return + } + generatedHandler(c) +} + // getGroupPlatform extracts the group platform from the API Key stored in context. func getGroupPlatform(c *gin.Context) string { apiKey, ok := middleware.GetAPIKeyFromContext(c) diff --git a/backend/internal/server/routes/gateway_codex_models_test.go b/backend/internal/server/routes/gateway_codex_models_test.go index 04a8b8fa67..0c6542a0fa 100644 --- a/backend/internal/server/routes/gateway_codex_models_test.go +++ b/backend/internal/server/routes/gateway_codex_models_test.go @@ -2,8 +2,12 @@ package routes import ( "net/http" + "net/http/httptest" "testing" + "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" ) @@ -22,3 +26,39 @@ func TestGatewayRoutesCodexModelsManifestPathIsRegistered(t *testing.T) { require.NotEmpty(t, registered["/models"], "GET /models should be registered") require.Equal(t, registered["/v1/models"], registered["/models"], "root alias should use the same platform-aware handler") } + +func TestDispatchCodexModelsGatewayKeepsOnlyOpenAIOnLiveManifestHandler(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + platform string + wantOpenAI bool + }{ + {platform: service.PlatformOpenAI, wantOpenAI: true}, + {platform: service.PlatformComposite}, + {platform: service.PlatformGrok}, + {platform: service.PlatformDeepseek}, + } + + for _, tt := range tests { + t.Run(tt.platform, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil) + c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{Platform: tt.platform}, + }) + called := "" + + dispatchCodexModelsGateway(c, + func(c *gin.Context) { called = "openai" }, + func(c *gin.Context) { called = "generated" }, + ) + + if tt.wantOpenAI { + require.Equal(t, "openai", called) + } else { + require.Equal(t, "generated", called) + } + }) + } +} diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index e7ce083ff5..4e304eca10 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -147,6 +147,9 @@ type AccountTestService struct { cfg *config.Config settingService *SettingService tlsFPProfileService *TLSFingerprintProfileService + modelMetadataRegistryMu sync.Mutex + modelMetadataRegistry map[string]modelsDevProvider + modelMetadataRegistryAt time.Time pluginManager *PluginManager agentIdentityTaskMu sync.Mutex agentIdentityWS agentIdentityWSConnectionInvalidator diff --git a/backend/internal/service/antigravity_gateway_compat.go b/backend/internal/service/antigravity_gateway_compat.go index 49ca076884..1d6d39e1aa 100644 --- a/backend/internal/service/antigravity_gateway_compat.go +++ b/backend/internal/service/antigravity_gateway_compat.go @@ -29,6 +29,8 @@ const ( AntigravityCredentialRejectedReason GatewayFailureReason = "antigravity_oauth_credential_rejected" ) +const antigravityCompatMaxTokens = 64000 + type antigravityCompatRequest struct { protocol antigravityCompatProtocol originalBody []byte @@ -158,7 +160,7 @@ func preserveChatCompletionTokenLimit(request *apicompat.ChatCompletionsRequest, limit = request.MaxCompletionTokens } if limit != nil && *limit > 0 { - claudeRequest.MaxTokens = *limit + claudeRequest.MaxTokens = min(*limit, antigravityCompatMaxTokens) } } diff --git a/backend/internal/service/antigravity_gateway_compat_test.go b/backend/internal/service/antigravity_gateway_compat_test.go index 3df0b0081a..a92b1654c2 100644 --- a/backend/internal/service/antigravity_gateway_compat_test.go +++ b/backend/internal/service/antigravity_gateway_compat_test.go @@ -12,6 +12,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -285,6 +286,21 @@ func TestAntigravityCompatPreservesChatTokenLimit(t *testing.T) { body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":8,"max_completion_tokens":13}`, want: 13, }, + { + name: "max_tokens at safe ceiling is preserved", + body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":64000}`, + want: 64000, + }, + { + name: "max_tokens above safe ceiling is clamped", + body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":64001}`, + want: 64000, + }, + { + name: "precedence applies before clamping", + body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":8,"max_completion_tokens":64001}`, + want: 64000, + }, } for _, tt := range tests { @@ -310,6 +326,27 @@ func TestAntigravityCompatPreservesChatTokenLimit(t *testing.T) { } } +func TestPreserveChatCompletionTokenLimitIgnoresAbsentAndNonPositiveValues(t *testing.T) { + tests := []struct { + name string + request apicompat.ChatCompletionsRequest + }{ + {name: "absent"}, + {name: "zero max_tokens", request: apicompat.ChatCompletionsRequest{MaxTokens: antigravityCompatIntPtr(0)}}, + {name: "negative max_completion_tokens takes precedence", request: apicompat.ChatCompletionsRequest{MaxTokens: antigravityCompatIntPtr(12), MaxCompletionTokens: antigravityCompatIntPtr(-1)}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + claudeRequest := &apicompat.AnthropicRequest{MaxTokens: 99} + preserveChatCompletionTokenLimit(&tt.request, claudeRequest) + require.Equal(t, 99, claudeRequest.MaxTokens) + }) + } +} + +func antigravityCompatIntPtr(v int) *int { return &v } + func TestAntigravityCompatRoutesByMappedModelFamily(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { diff --git a/backend/internal/service/auth_email_binding.go b/backend/internal/service/auth_email_binding.go index 4fe3c4e85b..733e76b7fd 100644 --- a/backend/internal/service/auth_email_binding.go +++ b/backend/internal/service/auth_email_binding.go @@ -56,12 +56,8 @@ func (s *AuthService) BindEmailIdentity( return nil, ErrPasswordIncorrect } - existingUser, err := s.userRepo.GetByEmail(ctx, normalizedEmail) - switch { - case err == nil && existingUser != nil && existingUser.ID != userID: - return nil, ErrEmailExists - case err != nil && !errors.Is(err, ErrUserNotFound): - return nil, ErrServiceUnavailable + if err := s.ensureEmailIdentityAvailableForUser(ctx, currentUser, normalizedEmail); err != nil { + return nil, err } hashedPassword, err := s.HashPassword(password) @@ -115,19 +111,16 @@ func (s *AuthService) SendEmailIdentityBindCode(ctx context.Context, userID int6 if s.emailService == nil { return ErrServiceUnavailable } - if _, err := s.userRepo.GetByID(ctx, userID); err != nil { + currentUser, err := s.userRepo.GetByID(ctx, userID) + if err != nil { if errors.Is(err, ErrUserNotFound) { return ErrUserNotFound } return ErrServiceUnavailable } - existingUser, err := s.userRepo.GetByEmail(ctx, normalizedEmail) - switch { - case err == nil && existingUser != nil && existingUser.ID != userID: - return ErrEmailExists - case err != nil && !errors.Is(err, ErrUserNotFound): - return ErrServiceUnavailable + if err := s.ensureEmailIdentityAvailableForUser(ctx, currentUser, normalizedEmail); err != nil { + return err } siteName := "Sub2API" @@ -137,6 +130,45 @@ func (s *AuthService) SendEmailIdentityBindCode(ctx context.Context, userID int6 return s.emailService.SendVerifyCode(ctx, normalizedEmail, siteName, firstEmailLocale(locale)) } +// ensureEmailIdentityAvailableForUser 在发码 / 提交换绑前做快速查重。 +// 精确地址或 provider alias 若已指向其他用户的收件箱则直接拒绝; +// 当前用户自己的收件箱允许继续,便于其更换自身的 alias 变体。 +func (s *AuthService) ensureEmailIdentityAvailableForUser( + ctx context.Context, + currentUser *User, + email string, +) error { + if currentUser == nil { + return ErrUserNotFound + } + + existingUser, err := s.userRepo.GetByEmail(ctx, email) + switch { + case err == nil: + if existingUser == nil || existingUser.ID == currentUser.ID { + break + } + return ErrEmailExists + case errors.Is(err, ErrUserNotFound): + // Continue to alias lookup below. + default: + return ErrServiceUnavailable + } + + if NormalizeEmailForAliasDedup(currentUser.Email) == NormalizeEmailForAliasDedup(email) { + return nil + } + + aliasExists, err := s.userRepo.ExistsByEmailAlias(ctx, email) + if err != nil { + return ErrServiceUnavailable + } + if aliasExists { + return ErrEmailExists + } + return nil +} + func normalizeEmailForIdentityBinding(email string) (string, error) { normalized := strings.ToLower(strings.TrimSpace(email)) if normalized == "" || len(normalized) > 255 { @@ -153,6 +185,12 @@ func hasBindableEmailIdentitySubject(email string) bool { return normalized != "" && !isReservedEmail(normalized) } +// emailIdentityAliasGuardRepository 是主邮箱替换所需的事务内原子仓储能力, +// 用于关闭服务层前置查重与实际写入之间的并发窗口。 +type emailIdentityAliasGuardRepository interface { + UpdateEmailWithAliasGuard(ctx context.Context, userID int64, email string, passwordHash string) error +} + func (s *AuthService) updateBoundEmailIdentityTx( ctx context.Context, currentUser *User, @@ -192,16 +230,15 @@ func (s *AuthService) updateBoundEmailIdentityWithClient( return ErrServiceUnavailable } - oldEmail := currentUser.Email - if _, err := client.User.UpdateOneID(currentUser.ID). - SetEmail(email). - SetPasswordHash(hashedPassword). - Save(ctx); err != nil { - if dbent.IsConstraintError(err) { - return ErrEmailExists - } + guard, ok := s.userRepo.(emailIdentityAliasGuardRepository) + if !ok { return ErrServiceUnavailable } + if err := guard.UpdateEmailWithAliasGuard(ctx, currentUser.ID, email, hashedPassword); err != nil { + return err + } + + oldEmail := currentUser.Email if err := replaceBoundEmailAuthIdentityWithClient(ctx, client, currentUser.ID, oldEmail, email, "auth_service_email_bind"); err != nil { if errors.Is(err, ErrEmailExists) { diff --git a/backend/internal/service/auth_service_email_bind_test.go b/backend/internal/service/auth_service_email_bind_test.go index 2d78862ef3..9bdca02f2d 100644 --- a/backend/internal/service/auth_service_email_bind_test.go +++ b/backend/internal/service/auth_service_email_bind_test.go @@ -6,6 +6,7 @@ import ( "context" "database/sql" "errors" + "fmt" "sync" "testing" "time" @@ -13,6 +14,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/authidentity" "github.com/Wei-Shaw/sub2api/ent/enttest" + dbuser "github.com/Wei-Shaw/sub2api/ent/user" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/repository" @@ -69,7 +71,8 @@ func newAuthServiceForEmailBindWithRefreshCache( ) (*service.AuthService, service.UserRepository, *dbent.Client) { t.Helper() - db, err := sql.Open("sqlite", "file:auth_service_email_bind?mode=memory&cache=shared") + dbName := fmt.Sprintf("file:auth_service_email_bind_%d?mode=memory&cache=shared", time.Now().UnixNano()) + db, err := sql.Open("sqlite", dbName) require.NoError(t, err) t.Cleanup(func() { _ = db.Close() }) @@ -214,6 +217,143 @@ func TestAuthServiceBindEmailIdentity_RejectsExistingEmailOnAnotherUser(t *testi require.Equal(t, 0, countProviderGrantRecords(t, client, sourceUser.ID, "email", "first_bind")) } +func TestAuthServiceBindEmailIdentity_RejectsAliasOfExistingEmailOnAnotherUser(t *testing.T) { + cache := &emailBindCacheStub{ + data: &service.VerificationCodeData{ + Code: "123456", + CreatedAt: time.Now().UTC().Add(-10 * time.Minute), + ExpiresAt: time.Now().UTC().Add(10 * time.Minute), + }, + } + svc, _, client := newAuthServiceForEmailBind(t, nil, cache, nil) + + ctx := context.Background() + sourceUser := createEmailBindTestUser( + t, + client, + "source-user"+service.OIDCConnectSyntheticEmailDomain, + "source-user", + "old-hash", + ) + createEmailBindTestUser(t, client, "zck.ioio123@gmail.com", "inbox-owner", "hash") + + err := svc.SendEmailIdentityBindCode(ctx, sourceUser.ID, "zckioio123+new@gmail.com") + require.ErrorIs(t, err, service.ErrEmailExists) + require.Empty(t, cache.setEmails) + + updatedUser, err := svc.BindEmailIdentity( + ctx, + sourceUser.ID, + "zckioio123+new@gmail.com", + "123456", + "new-password", + ) + require.ErrorIs(t, err, service.ErrEmailExists) + require.Nil(t, updatedUser) + + storedUser, err := client.User.Get(ctx, sourceUser.ID) + require.NoError(t, err) + require.Equal(t, "source-user"+service.OIDCConnectSyntheticEmailDomain, storedUser.Email) + require.Equal(t, "old-hash", storedUser.PasswordHash) +} + +func TestAuthServiceBindEmailIdentity_AllowsOnlyOneConcurrentAliasVariant(t *testing.T) { + cache := &emailBindCacheStub{ + data: &service.VerificationCodeData{ + Code: "123456", + CreatedAt: time.Now().UTC(), + ExpiresAt: time.Now().UTC().Add(10 * time.Minute), + }, + } + svc, _, client := newAuthServiceForEmailBind(t, nil, cache, nil) + + ctx := context.Background() + unique := fmt.Sprintf("%d", time.Now().UnixNano()) + first := createEmailBindTestUser( + t, + client, + "first-"+unique+service.OIDCConnectSyntheticEmailDomain, + "first-"+unique, + "old-hash", + ) + second := createEmailBindTestUser( + t, + client, + "second-"+unique+service.OIDCConnectSyntheticEmailDomain, + "second-"+unique, + "old-hash", + ) + + start := make(chan struct{}) + results := make(chan error, 2) + go func() { + <-start + _, err := svc.BindEmailIdentity(ctx, first.ID, "inbox-"+unique+"+one@gmail.com", "123456", "new-password") + results <- err + }() + go func() { + <-start + _, err := svc.BindEmailIdentity(ctx, second.ID, "inbox-"+unique+"+two@gmail.com", "123456", "new-password") + results <- err + }() + close(start) + + var successes, conflicts int + for range 2 { + err := <-results + switch { + case err == nil: + successes++ + case errors.Is(err, service.ErrEmailExists): + conflicts++ + default: + t.Fatalf("unexpected bind error: %v", err) + } + } + require.Equal(t, 1, successes) + require.Equal(t, 1, conflicts) + + boundCount, err := client.User.Query(). + Where(dbuser.EmailIn( + "inbox-"+unique+"+one@gmail.com", + "inbox-"+unique+"+two@gmail.com", + )). + Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, boundCount) +} + +func TestAuthServiceBindEmailIdentity_RejectsNewAliasWhenAnotherUserSharesCurrentUserInbox(t *testing.T) { + cache := &emailBindCacheStub{ + data: &service.VerificationCodeData{ + Code: "123456", + CreatedAt: time.Now().UTC(), + ExpiresAt: time.Now().UTC().Add(10 * time.Minute), + }, + } + svc, _, client := newAuthServiceForEmailBind(t, nil, cache, nil) + + ctx := context.Background() + hashedPassword, err := svc.HashPassword("current-password") + require.NoError(t, err) + currentUser := createEmailBindTestUser(t, client, "inbox+own@gmail.com", "current", hashedPassword) + createEmailBindTestUser(t, client, "inbox+legacy@gmail.com", "legacy", "hash") + + updatedUser, err := svc.BindEmailIdentity( + ctx, + currentUser.ID, + "inbox+new@gmail.com", + "123456", + "current-password", + ) + require.ErrorIs(t, err, service.ErrEmailExists) + require.Nil(t, updatedUser) + + storedUser, err := client.User.Get(ctx, currentUser.ID) + require.NoError(t, err) + require.Equal(t, "inbox+own@gmail.com", storedUser.Email) +} + func TestAuthServiceBindEmailIdentity_RollsBackWhenFirstBindDefaultsFail(t *testing.T) { assigner := &flakyEmailBindDefaultSubAssignerStub{err: errors.New("temporary assign failure")} cache := &emailBindCacheStub{ diff --git a/backend/internal/service/composite_model_route.go b/backend/internal/service/composite_model_route.go index a008a3e9fa..1201dad5e4 100644 --- a/backend/internal/service/composite_model_route.go +++ b/backend/internal/service/composite_model_route.go @@ -23,8 +23,19 @@ const ( CompositeRouteSourceExplicit = "route" CompositeRouteSourceDetector = "detector" + CompositeRouteSourceAccount = "account_model" ) +// CompositeModelOwnership identifies the concrete provider that exposes a +// public model through an account-level exact mapping. +type CompositeModelOwnership struct { + TargetPlatform string + Matched bool + Ambiguous bool +} + +type CompositeModelOwnershipResolver func(context.Context, int64, string) (CompositeModelOwnership, error) + var ( ErrCompositeRouteNotFound = infraerrors.NotFound("COMPOSITE_ROUTE_NOT_FOUND", "composite route not found") ErrCompositeRouteExists = infraerrors.Conflict("COMPOSITE_ROUTE_EXISTS", "composite route already exists") diff --git a/backend/internal/service/composite_platform_test.go b/backend/internal/service/composite_platform_test.go index 5373a7fc6f..bff5928024 100644 --- a/backend/internal/service/composite_platform_test.go +++ b/backend/internal/service/composite_platform_test.go @@ -8,6 +8,150 @@ import ( "github.com/stretchr/testify/require" ) +type compositeOwnershipAccountRepo struct { + AccountRepository + accounts []Account +} + +func (r *compositeOwnershipAccountRepo) ListSchedulableByGroupID(context.Context, int64) ([]Account, error) { + return r.accounts, nil +} + +// Scenario: 唯一平台的精确别名可路由 +func TestResolveCompositeModelOwnershipKeepsProviderAccountsIsolated(t *testing.T) { + groupID := int64(7) + repo := &compositeOwnershipAccountRepo{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-public": "gpt-5"}, + }, + }, + { + ID: 2, + Platform: PlatformDeepseek, + Credentials: map[string]any{ + "model_mapping": map[string]any{"reasoning-alias": "deepseek-v4-pro"}, + }, + }, + }, + } + svc := &GatewayService{accountRepo: repo} + + deepSeekOwnership, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "reasoning-alias") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformDeepseek, Matched: true}, deepSeekOwnership) + + openAIOwnership, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "gpt-public") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformOpenAI, Matched: true}, openAIOwnership) +} + +// Scenario: 通配符和空映射不声明所有权 +func TestResolveCompositeModelOwnershipRequiresNonEmptyExactMappings(t *testing.T) { + groupID := int64(7) + repo := &compositeOwnershipAccountRepo{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"*": "gpt-5", "gpt-*": "gpt-5", "empty-alias": ""}, + }, + }, + { + ID: 2, + Platform: PlatformGrok, + Credentials: map[string]any{ + "model_mapping": map[string]any{"grok-public": "grok-4"}, + }, + }, + }, + } + svc := &GatewayService{accountRepo: repo} + + for _, model := range []string{"gpt-5", "empty-alias", "unknown-alias"} { + ownership, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, model) + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{}, ownership, "model=%s", model) + } + + ownership, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "grok-public") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformGrok, Matched: true}, ownership) +} + +func TestResolveCompositeModelOwnershipAllowsSamePlatformAndRejectsCrossPlatformAliases(t *testing.T) { + groupID := int64(7) + repo := &compositeOwnershipAccountRepo{ + accounts: []Account{ + {ID: 1, Platform: PlatformOpenAI, Credentials: map[string]any{"model_mapping": map[string]any{"shared-openai": "gpt-5", "ambiguous": "gpt-5"}}}, + {ID: 2, Platform: PlatformOpenAI, Credentials: map[string]any{"model_mapping": map[string]any{"shared-openai": "gpt-5.1"}}}, + {ID: 3, Platform: PlatformDeepseek, Credentials: map[string]any{"model_mapping": map[string]any{"ambiguous": "deepseek-v4-pro"}}}, + }, + } + svc := &GatewayService{accountRepo: repo} + + samePlatform, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "shared-openai") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformOpenAI, Matched: true}, samePlatform) + + ambiguous, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "ambiguous") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{Ambiguous: true}, ambiguous) +} + +func TestNewGatewayServiceWiresCompositeModelOwnershipResolver(t *testing.T) { + groupID := int64(7) + repo := &compositeOwnershipAccountRepo{ + accounts: []Account{{ + ID: 1, + Platform: PlatformDeepseek, + Credentials: map[string]any{"model_mapping": map[string]any{"reasoning-alias": "deepseek-v4-pro"}}, + }}, + } + resolver := NewCompositeRouteResolver(nil) + svc := NewGatewayService( + repo, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + resolver, + nil, + nil, + ) + require.Same(t, resolver, svc.compositeResolver) + + decision, err := resolver.Resolve(context.Background(), groupID, "reasoning-alias", CompositeRouteEndpointResponses) + require.NoError(t, err) + require.True(t, decision.Matched) + require.Equal(t, CompositeRouteSourceAccount, decision.Source) + require.Equal(t, PlatformDeepseek, decision.TargetPlatform) +} + func TestDetectModelPlatform(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/composite_route_resolver.go b/backend/internal/service/composite_route_resolver.go index 3d69f13553..f2e1928ace 100644 --- a/backend/internal/service/composite_route_resolver.go +++ b/backend/internal/service/composite_route_resolver.go @@ -8,13 +8,20 @@ import ( ) type CompositeRouteResolver struct { - repo CompositeModelRouteRepository + repo CompositeModelRouteRepository + modelOwnershipResolver CompositeModelOwnershipResolver } func NewCompositeRouteResolver(repo CompositeModelRouteRepository) *CompositeRouteResolver { return &CompositeRouteResolver{repo: repo} } +func (r *CompositeRouteResolver) SetModelOwnershipResolver(resolver CompositeModelOwnershipResolver) { + if r != nil { + r.modelOwnershipResolver = resolver + } +} + func (r *CompositeRouteResolver) Resolve(ctx context.Context, groupID int64, model, endpoint string) (CompositeRouteDecision, error) { model = strings.TrimSpace(model) endpoint = normalizeCompositeRouteEndpoint(endpoint) @@ -51,6 +58,35 @@ func (r *CompositeRouteResolver) Resolve(ctx context.Context, groupID int64, mod } } + if r != nil && r.modelOwnershipResolver != nil && groupID > 0 { + ownership, err := r.modelOwnershipResolver(ctx, groupID, model) + if err != nil { + // A recognizable model can still use the existing detector when the + // account catalog is temporarily unavailable. Unknown aliases cannot. + if _, detectable := DetectModelPlatform(model); !detectable { + return decision, fmt.Errorf("resolve account model ownership: %w", err) + } + } else if ownership.Ambiguous { + decision.Reason = "model is exposed by multiple provider platforms" + return decision, nil + } else if ownership.Matched { + platform := strings.TrimSpace(ownership.TargetPlatform) + if !isConcreteRequestPlatform(platform) { + decision.Reason = "account model ownership has no concrete target platform" + return decision, nil + } + return CompositeRouteDecision{ + Matched: true, + Source: CompositeRouteSourceAccount, + GroupID: groupID, + PublicModel: model, + TargetPlatform: platform, + UpstreamModel: model, + Endpoint: endpoint, + }, nil + } + } + if platform, ok := DetectModelPlatform(model); ok { return CompositeRouteDecision{ Matched: true, diff --git a/backend/internal/service/composite_route_resolver_test.go b/backend/internal/service/composite_route_resolver_test.go index 60e4bd0b82..474b29387b 100644 --- a/backend/internal/service/composite_route_resolver_test.go +++ b/backend/internal/service/composite_route_resolver_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "testing" "github.com/stretchr/testify/require" @@ -69,6 +70,98 @@ func TestCompositeRouteResolverExplicitExactRouteRewritesModel(t *testing.T) { require.Equal(t, int64(10), decision.Route.ID) } +// Scenario: 唯一平台的精确别名可路由 +func TestCompositeRouteResolverUsesAccountModelOwnershipForUnprefixedAlias(t *testing.T) { + resolver := NewCompositeRouteResolver(nil) + resolver.SetModelOwnershipResolver(func(_ context.Context, groupID int64, model string) (CompositeModelOwnership, error) { + require.Equal(t, int64(7), groupID) + require.Equal(t, "reasoning-alias", model) + return CompositeModelOwnership{TargetPlatform: PlatformDeepseek, Matched: true}, nil + }) + + decision, err := resolver.Resolve(context.Background(), 7, "reasoning-alias", CompositeRouteEndpointChatCompletions) + + require.NoError(t, err) + require.True(t, decision.Matched) + require.Equal(t, CompositeRouteSourceAccount, decision.Source) + require.Equal(t, PlatformDeepseek, decision.TargetPlatform) + require.Equal(t, "reasoning-alias", decision.UpstreamModel) +} + +func TestCompositeRouteResolverAccountOwnershipOverridesBuiltInDetector(t *testing.T) { + resolver := NewCompositeRouteResolver(nil) + resolver.SetModelOwnershipResolver(func(context.Context, int64, string) (CompositeModelOwnership, error) { + return CompositeModelOwnership{TargetPlatform: PlatformDeepseek, Matched: true}, nil + }) + + decision, err := resolver.Resolve(context.Background(), 7, "gpt-5", CompositeRouteEndpointResponses) + + require.NoError(t, err) + require.True(t, decision.Matched) + require.Equal(t, CompositeRouteSourceAccount, decision.Source) + require.Equal(t, PlatformDeepseek, decision.TargetPlatform) +} + +// Scenario: 显式路由保持最高优先级 +func TestCompositeRouteResolverExplicitRouteBeatsAccountOwnership(t *testing.T) { + resolver := NewCompositeRouteResolver(compositeRouteRepoStub{ + routes: []CompositeModelRoute{{ + ID: 10, + GroupID: 7, + PublicModel: "reasoning-alias", + MatchType: CompositeRouteMatchExact, + TargetPlatform: PlatformOpenAI, + UpstreamModel: "gpt-5", + Endpoint: CompositeRouteEndpointAny, + Enabled: true, + }}, + }) + resolver.SetModelOwnershipResolver(func(context.Context, int64, string) (CompositeModelOwnership, error) { + return CompositeModelOwnership{TargetPlatform: PlatformDeepseek, Matched: true}, nil + }) + + decision, err := resolver.Resolve(context.Background(), 7, "reasoning-alias", CompositeRouteEndpointChatCompletions) + + require.NoError(t, err) + require.True(t, decision.Matched) + require.Equal(t, CompositeRouteSourceExplicit, decision.Source) + require.Equal(t, PlatformOpenAI, decision.TargetPlatform) + require.Equal(t, "gpt-5", decision.UpstreamModel) +} + +// Scenario: 跨平台同名别名不被猜测 +func TestCompositeRouteResolverDoesNotGuessAmbiguousAccountOwnership(t *testing.T) { + resolver := NewCompositeRouteResolver(nil) + resolver.SetModelOwnershipResolver(func(context.Context, int64, string) (CompositeModelOwnership, error) { + return CompositeModelOwnership{Ambiguous: true}, nil + }) + + decision, err := resolver.Resolve(context.Background(), 7, "shared-alias", CompositeRouteEndpointChatCompletions) + + require.NoError(t, err) + require.False(t, decision.Matched) + require.Empty(t, decision.TargetPlatform) + require.Equal(t, "model is exposed by multiple provider platforms", decision.Reason) +} + +func TestCompositeRouteResolverOwnershipLookupErrorFallsBackOnlyForDetectableModels(t *testing.T) { + lookupErr := errors.New("account catalog unavailable") + resolver := NewCompositeRouteResolver(nil) + resolver.SetModelOwnershipResolver(func(context.Context, int64, string) (CompositeModelOwnership, error) { + return CompositeModelOwnership{}, lookupErr + }) + + detected, err := resolver.Resolve(context.Background(), 7, "gpt-5", CompositeRouteEndpointResponses) + require.NoError(t, err) + require.True(t, detected.Matched) + require.Equal(t, CompositeRouteSourceDetector, detected.Source) + require.Equal(t, PlatformOpenAI, detected.TargetPlatform) + + unknown, err := resolver.Resolve(context.Background(), 7, "company-model", CompositeRouteEndpointResponses) + require.ErrorIs(t, err, lookupErr) + require.False(t, unknown.Matched) +} + func TestCompositeRouteResolverPrefersEndpointSpecificLongestPrefix(t *testing.T) { resolver := NewCompositeRouteResolver(compositeRouteRepoStub{ routes: []CompositeModelRoute{ diff --git a/backend/internal/service/gateway_hotpath_optimization_test.go b/backend/internal/service/gateway_hotpath_optimization_test.go index 1bdf48bd41..8855b13e52 100644 --- a/backend/internal/service/gateway_hotpath_optimization_test.go +++ b/backend/internal/service/gateway_hotpath_optimization_test.go @@ -564,6 +564,46 @@ func TestGetAvailableModels_UsesShortCacheAndSupportsInvalidation(t *testing.T) require.Equal(t, int64(2), store) } +// Scenario: 账号模型变更会失效所属平台缓存 +func TestResolveCompositeModelOwnershipUsesModelsCacheInvalidation(t *testing.T) { + groupID := int64(9) + repo := &modelsListAccountRepoStub{ + byGroup: map[int64][]Account{ + groupID: {{ + ID: 1, + Platform: PlatformDeepseek, + Credentials: map[string]any{"model_mapping": map[string]any{"company-model": "deepseek-v4-pro"}}, + }}, + }, + } + svc := &GatewayService{ + accountRepo: repo, + modelsListCache: gocache.New(time.Minute, time.Minute), + modelsListCacheTTL: time.Minute, + } + + first, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "company-model") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformDeepseek, Matched: true}, first) + require.Equal(t, int64(1), repo.listByGroupCalls.Load()) + + repo.byGroup[groupID] = []Account{{ + ID: 2, + Platform: PlatformOpenAI, + Credentials: map[string]any{"model_mapping": map[string]any{"company-model": "gpt-5"}}, + }} + cached, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "company-model") + require.NoError(t, err) + require.Equal(t, first, cached) + require.Equal(t, int64(1), repo.listByGroupCalls.Load()) + + svc.InvalidateAvailableModelsCache(&groupID, PlatformDeepseek) + refreshed, err := svc.resolveCompositeModelOwnership(context.Background(), groupID, "company-model") + require.NoError(t, err) + require.Equal(t, CompositeModelOwnership{TargetPlatform: PlatformOpenAI, Matched: true}, refreshed) + require.Equal(t, int64(2), repo.listByGroupCalls.Load()) +} + func TestGetAvailableModels_ErrorAndGlobalListBranches(t *testing.T) { resetGatewayHotpathStatsForTest() diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index 664b6bf08b..a8834bc403 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -395,6 +395,61 @@ func TestGatewayService_SelectAccountForModelWithPlatform_Anthropic(t *testing.T require.Equal(t, PlatformAnthropic, acc.Platform, "应只返回 anthropic 平台账户") } +// Scenario: account-owned Composite aliases are scheduled only to accounts that declare the exact mapping. +func TestGatewayService_SelectAccountForModelWithExclusions_CompositeAliasRequiresOwningAccount(t *testing.T) { + groupID := int64(77) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Priority: 1, + Status: StatusActive, + Schedulable: true, + AccountGroups: []AccountGroup{{GroupID: groupID}}, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "model_mapping": map[string]any{"reasoning-alias": "claude-opus-4-8"}, + }, + AccountGroups: []AccountGroup{{GroupID: groupID}}, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + group := &Group{ID: groupID, Platform: PlatformComposite, Status: StatusActive, Hydrated: true} + svc := &GatewayService{ + accountRepo: repo, + groupRepo: &mockGroupRepoForGateway{groups: map[int64]*Group{groupID: group}}, + cfg: testConfig(), + } + ctx := WithCompositeRouteDecision(context.Background(), CompositeRouteDecision{ + Matched: true, + Source: CompositeRouteSourceAccount, + GroupID: groupID, + PublicModel: "reasoning-alias", + TargetPlatform: PlatformAnthropic, + UpstreamModel: "reasoning-alias", + Endpoint: CompositeRouteEndpointResponses, + }) + + account, err := svc.SelectAccountForModelWithExclusions(ctx, &groupID, "", "reasoning-alias", nil) + require.NoError(t, err) + require.NotNil(t, account) + require.Equal(t, int64(2), account.ID) +} + // TestGatewayService_SelectAccountForModelWithPlatform_Antigravity 测试 antigravity 单平台选择 func TestGatewayService_SelectAccountForModelWithPlatform_Antigravity(t *testing.T) { ctx := context.Background() diff --git a/backend/internal/service/gateway_scheduling.go b/backend/internal/service/gateway_scheduling.go index be63f804f5..977a28a99d 100644 --- a/backend/internal/service/gateway_scheduling.go +++ b/backend/internal/service/gateway_scheduling.go @@ -2532,6 +2532,11 @@ func summarizeSelectionFailureStats(stats selectionFailureStats) string { // isModelSupportedByAccountWithContext 根据账户平台检查模型支持(带 context) // 对于 Antigravity 平台,会先获取映射后的最终模型名(包括 thinking 后缀)再检查支持 func (s *GatewayService) isModelSupportedByAccountWithContext(ctx context.Context, account *Account, requestedModel string) bool { + if source, ok := CompositeRouteSourceFromContext(ctx); ok && source == CompositeRouteSourceAccount { + if publicModel, modelOK := RequestedPublicModelFromContext(ctx); modelOK && !explicitModelMappingClaims(*account, publicModel) { + return false + } + } if account.Platform == PlatformAntigravity { if strings.TrimSpace(requestedModel) == "" { return true diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 37bab2f97d..8910cec117 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -71,8 +71,9 @@ const ( ) const ( - cacheTTLTarget5m = "5m" - cacheTTLTarget1h = "1h" + cacheTTLTarget5m = "5m" + cacheTTLTarget1h = "1h" + compositeModelOwnershipCachePrefix = "composite-owner|" ) // ForceCacheBillingContextKey 强制缓存计费上下文键 @@ -528,6 +529,10 @@ func modelsListCacheKey(groupID *int64, platform string) string { return fmt.Sprintf("%d|%s", derefGroupID(groupID), strings.TrimSpace(platform)) } +func compositeModelOwnershipCacheKey(groupID int64, model string) string { + return fmt.Sprintf("%s%d|%s", compositeModelOwnershipCachePrefix, groupID, strings.TrimSpace(model)) +} + func prefetchedStickyGroupIDFromContext(ctx context.Context) (int64, bool) { return PrefetchedStickyGroupIDFromContext(ctx) } @@ -858,6 +863,9 @@ func NewGatewayService( balanceNotifyService: balanceNotifyService, userPlatformQuotaRepo: userPlatformQuotaRepo, } + if compositeResolver != nil { + compositeResolver.SetModelOwnershipResolver(svc.resolveCompositeModelOwnership) + } svc.userGroupRateResolver = newUserGroupRateResolver( userGroupRateRepo, svc.userGroupRateCache, @@ -1447,6 +1455,59 @@ func (s *GatewayService) GetAvailableModels(ctx context.Context, groupID *int64, return cloneStringSlice(models) } +func (s *GatewayService) resolveCompositeModelOwnership(ctx context.Context, groupID int64, model string) (CompositeModelOwnership, error) { + model = strings.TrimSpace(model) + if s == nil || s.accountRepo == nil || groupID <= 0 || model == "" { + return CompositeModelOwnership{}, nil + } + + cacheKey := compositeModelOwnershipCacheKey(groupID, model) + if s.modelsListCache != nil { + if cached, found := s.modelsListCache.Get(cacheKey); found { + if ownership, ok := cached.(CompositeModelOwnership); ok { + return ownership, nil + } + } + } + + accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, groupID) + if err != nil { + return CompositeModelOwnership{}, err + } + + platforms := make(map[string]struct{}) + for _, account := range accounts { + platform := strings.TrimSpace(account.Platform) + if !isConcreteRequestPlatform(platform) || !explicitModelMappingClaims(account, model) { + continue + } + platforms[platform] = struct{}{} + } + + ownership := CompositeModelOwnership{} + if len(platforms) == 1 { + for platform := range platforms { + ownership.TargetPlatform = platform + } + ownership.Matched = true + } else if len(platforms) > 1 { + ownership.Ambiguous = true + } + + if s.modelsListCache != nil { + s.modelsListCache.Set(cacheKey, ownership, s.modelsListCacheTTL) + } + return ownership, nil +} + +func explicitModelMappingClaims(account Account, model string) bool { + if account.Credentials == nil || model == "" { + return false + } + mapped, ok := stringMappingFromRaw(account.Credentials["model_mapping"])[model] + return ok && strings.TrimSpace(mapped) != "" +} + // GetSchedulablePlatforms returns the concrete platforms that currently have // schedulable accounts in the target group. func (s *GatewayService) GetSchedulablePlatforms(ctx context.Context, groupID *int64) map[string]struct{} { @@ -1479,6 +1540,7 @@ func (s *GatewayService) InvalidateAvailableModelsCache(groupID *int64, platform if s == nil || s.modelsListCache == nil { return } + s.invalidateCompositeModelOwnershipCache(groupID) normalizedPlatform := strings.TrimSpace(platform) // 完整匹配时精准失效;否则按维度批量失效。 @@ -1507,6 +1569,26 @@ func (s *GatewayService) InvalidateAvailableModelsCache(groupID *int64, platform } } +func (s *GatewayService) invalidateCompositeModelOwnershipCache(groupID *int64) { + for key := range s.modelsListCache.Items() { + if !strings.HasPrefix(key, compositeModelOwnershipCachePrefix) { + continue + } + if groupID == nil { + s.modelsListCache.Delete(key) + continue + } + parts := strings.SplitN(strings.TrimPrefix(key, compositeModelOwnershipCachePrefix), "|", 2) + if len(parts) != 2 { + continue + } + cachedGroupID, err := strconv.ParseInt(parts[0], 10, 64) + if err == nil && cachedGroupID == *groupID { + s.modelsListCache.Delete(key) + } + } +} + const debugGatewayBodyDefaultFilename = "gateway_debug.log" // initDebugGatewayBodyFile 初始化网关调试日志文件。 diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index 1d8da95913..7b59c28fd4 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -28,6 +28,49 @@ type OpenAIOAuth429FailoverState struct { grokOAuth429FollowupPending bool } +type openAIOAuth429Disposition uint8 + +const ( + openAIOAuth429Transient openAIOAuth429Disposition = iota + openAIOAuth429Quota5h + openAIOAuth429Quota7d + openAIOAuth429QuotaReset +) + +// classifyOpenAIOAuth429 区分账号配额耗尽信号与普通瞬时 429。明确窗口达到 +// 100% 时以该窗口为准;没有 100% 标记但包含重置头时,沿用 v179 的兼容语义, +// 仍视为配额限流信号。 +func classifyOpenAIOAuth429(headers http.Header, responseBody []byte) (openAIOAuth429Disposition, *time.Time) { + if snapshot := ParseCodexRateLimitHeaders(headers); snapshot != nil { + if normalized := snapshot.Normalize(); normalized != nil { + if normalized.Used7dPercent != nil && *normalized.Used7dPercent >= 100 { + if normalized.Reset7dSeconds != nil { + now := time.Now() + resetAt := now.Add(time.Duration(*normalized.Reset7dSeconds) * time.Second) + return openAIOAuth429Quota7d, &resetAt + } + return openAIOAuth429Quota7d, nil + } + if normalized.Used5hPercent != nil && *normalized.Used5hPercent >= 100 { + if normalized.Reset5hSeconds != nil { + now := time.Now() + resetAt := now.Add(time.Duration(*normalized.Reset5hSeconds) * time.Second) + return openAIOAuth429Quota5h, &resetAt + } + return openAIOAuth429Quota5h, nil + } + } + } + if resetAt := calculateOpenAI429ResetTime(headers); resetAt != nil { + return openAIOAuth429QuotaReset, resetAt + } + if resetUnix := parseOpenAIRateLimitResetTime(responseBody); resetUnix != nil { + resetAt := time.Unix(*resetUnix, 0) + return openAIOAuth429QuotaReset, &resetAt + } + return openAIOAuth429Transient, nil +} + func openAIAccountStateContext(ctx context.Context) (context.Context, context.CancelFunc) { base := context.Background() if ctx != nil { @@ -90,6 +133,17 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont return false } + // Self-built images requests always carry a matching image_generation tool, so a + // "tool choice not found in 'tools'" 400 means upstream revoked this account's + // image capability. Gated on the self-built marker: passthrough clients control + // their own tools/tool_choice and could otherwise poison a healthy account. + if isOpenAIImagesSelfBuiltRequest(ctx) && isOpenAIImageCapabilityLossError(statusCode, responseBody) { + if s != nil && s.rateLimitService != nil { + _ = s.rateLimitService.HandleOpenAIImageCapabilityLoss(stateCtx, account, statusCode, responseBody) + } + return false + } + if s == nil || account == nil { return false } @@ -165,19 +219,16 @@ func (s *OpenAIGatewayService) markOpenAIOAuth429RateLimited(ctx context.Context return } s.recordOpenAIOAuth429() - if s.openAIOAuth429RetryWindowActive(account) { + disposition, resetAt := classifyOpenAIOAuth429(headers, responseBody) + if disposition == openAIOAuth429Transient && s.openAIOAuth429RetryWindowActive(account) { return } cooldownUntil := time.Now().Add(openAIOAuth429FallbackCooldown) - if s.rateLimitService != nil { - if resetAt := s.rateLimitService.calculateOpenAI429ResetTime(headers); resetAt != nil && resetAt.After(time.Now()) { - cooldownUntil = *resetAt - } else if resetUnix := parseOpenAIRateLimitResetTime(responseBody); resetUnix != nil { - if resetAt := time.Unix(*resetUnix, 0); resetAt.After(time.Now()) { - cooldownUntil = resetAt - } - } else if cooldown, ok := s.rateLimitService.get429FallbackCooldown(ctx, account); ok && cooldown > 0 { + if resetAt != nil && resetAt.After(time.Now()) { + cooldownUntil = *resetAt + } else if s.rateLimitService != nil { + if cooldown, ok := s.rateLimitService.get429FallbackCooldown(ctx, account); ok && cooldown > 0 { cooldownUntil = time.Now().Add(cooldown) } } @@ -186,9 +237,17 @@ func (s *OpenAIGatewayService) markOpenAIOAuth429RateLimited(ctx context.Context } func (s *OpenAIGatewayService) shouldRetryOpenAIOAuth429OnSameAccount(account *Account, statusCode int, shouldDisable bool) bool { + return s.shouldRetryOpenAIOAuth429OnSameAccountWithResponse(account, statusCode, shouldDisable, nil, nil) +} + +func (s *OpenAIGatewayService) shouldRetryOpenAIOAuth429OnSameAccountWithResponse(account *Account, statusCode int, shouldDisable bool, headers http.Header, responseBody []byte) bool { if shouldDisable || statusCode != http.StatusTooManyRequests || !isOpenAIOAuthAccount(account) || account.IsShadow() { return false } + disposition, _ := classifyOpenAIOAuth429(headers, responseBody) + if disposition != openAIOAuth429Transient { + return false + } // markOpenAIOAuth429RateLimited parks the account once the window expires. // Do not accidentally create a fresh window after that transition. if s.isOpenAIAccountRuntimeBlocked(account) { @@ -199,10 +258,14 @@ func (s *OpenAIGatewayService) shouldRetryOpenAIOAuth429OnSameAccount(account *A // ShouldRetryOpenAIOAuth429 lets RateLimitService defer persistent account // cooldown until the gateway's same-account retry window is exhausted. -func (s *OpenAIGatewayService) ShouldRetryOpenAIOAuth429(account *Account, _ http.Header, _ []byte) bool { +func (s *OpenAIGatewayService) ShouldRetryOpenAIOAuth429(account *Account, headers http.Header, responseBody []byte) bool { if s == nil || !isOpenAIOAuthAccount(account) || account.IsShadow() || s.isOpenAIAccountRuntimeBlocked(account) { return false } + disposition, _ := classifyOpenAIOAuth429(headers, responseBody) + if disposition != openAIOAuth429Transient { + return false + } return s.openAIOAuth429RetryWindowActive(account) } diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go index fcda6d6a05..89d8c5b5ec 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go @@ -14,7 +14,7 @@ import ( ) type oauth429RateLimitRepo struct { - AccountRepository + mockAccountRepoForGemini setRateLimitedCalls int lastRateLimitedUntil time.Time } @@ -62,6 +62,38 @@ func TestOpenAI429FastPath_BlocksOAuthOnlyAfterRetryWindow(t *testing.T) { require.False(t, svc.shouldRetryOpenAIOAuth429OnSameAccount(account, http.StatusTooManyRequests, false)) } +func TestOpenAI429FastPath_BlocksOAuthImmediatelyWhenSevenDayQuotaIsExhausted(t *testing.T) { + repo := &oauth429RateLimitRepo{} + rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + svc := &OpenAIGatewayService{rateLimitService: rateLimits} + rateLimits.SetAccountRuntimeBlocker(svc) + account := &Account{ID: 423, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + headers := http.Header{} + headers.Set("x-codex-primary-used-percent", "100") + headers.Set("x-codex-primary-reset-after-seconds", "604800") + headers.Set("x-codex-primary-window-minutes", "10080") + headers.Set("x-codex-secondary-used-percent", "20") + headers.Set("x-codex-secondary-reset-after-seconds", "3600") + headers.Set("x-codex-secondary-window-minutes", "300") + + shouldDisable := svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, headers, []byte(`{"error":{"type":"rate_limit_error","code":"rate_limit_exceeded"}}`)) + + require.False(t, shouldDisable) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, 1, repo.setRateLimitedCalls) + require.Greater(t, time.Until(repo.lastRateLimitedUntil), 6*24*time.Hour) + require.False(t, svc.ShouldRetryOpenAIOAuth429(account, headers, nil)) +} + +func TestOpenAI429FastPath_RetriesOAuthWhenNoQuotaSignalExists(t *testing.T) { + svc := &OpenAIGatewayService{} + account := &Account{ID: 424, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + headers := http.Header{"Retry-After": []string{"1"}} + + require.True(t, svc.ShouldRetryOpenAIOAuth429(account, headers, []byte(`{"error":{"type":"rate_limit_error","message":"try again"}}`))) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + func TestOpenAIStream429IgnoresSuccessfulQuotaSnapshotHeaders(t *testing.T) { repo := &oauth429RateLimitRepo{} rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) @@ -164,7 +196,7 @@ func TestOpenAI429FastPath_SkipsSparkShadow(t *testing.T) { svc.markOpenAIOAuth429RateLimited(context.Background(), normal, headers, nil) require.False(t, svc.isOpenAIAccountRuntimeBlocked(shadow), "spark shadow must not be runtime-blocked by /responses global 429") - require.False(t, svc.isOpenAIAccountRuntimeBlocked(normal), "normal OpenAI OAuth account stays schedulable during its retry window") + require.True(t, svc.isOpenAIAccountRuntimeBlocked(normal), "normal OpenAI OAuth account with an exhausted 5h window must be paused") } func TestOpenAIRuntimeBlock_AppliesToOpenAIAPIKeyWhenRateLimitServiceStopsScheduling(t *testing.T) { diff --git a/backend/internal/service/openai_codex_model_metadata.go b/backend/internal/service/openai_codex_model_metadata.go new file mode 100644 index 0000000000..c779a05ab9 --- /dev/null +++ b/backend/internal/service/openai_codex_model_metadata.go @@ -0,0 +1,313 @@ +package service + +import "strings" + +func groupCodexModelMetadata( + platform string, + modelID string, + accounts []Account, + compositeRoutes []CompositeModelRoute, + compositeRoutesAvailable bool, +) (codexModelMetadataOverride, bool) { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + return codexModelMetadataOverride{}, false + } + upstreamModel := modelID + if platform == PlatformComposite { + var resolved bool + platform, upstreamModel, resolved = resolveCodexCompositeModelTarget( + modelID, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ) + if !resolved { + if codexExplicitModelTargetsConflict(accounts, modelID) { + return codexModelMetadataOverride{ + reasoningConflict: true, + inputModalitiesConflict: true, + }, true + } + return codexModelMetadataOverride{}, false + } + } + if !isConcreteRequestPlatform(platform) { + return codexModelMetadataOverride{}, false + } + + explicitClaims := false + if upstreamModel == modelID { + for _, account := range accounts { + if account.Platform == platform && codexExplicitModelMappingClaims(account, modelID) { + explicitClaims = true + break + } + } + } + explicitTargetsConflict := explicitClaims && codexExplicitModelTargetsConflictForPlatform(accounts, platform, modelID) + publicAlias := upstreamModel != modelID + candidates := make([]UpstreamModelMetadata, 0) + for i := range accounts { + account := &accounts[i] + if account.Platform != platform { + continue + } + var lookupModel string + if explicitClaims { + if !codexExplicitModelMappingClaims(*account, modelID) { + continue + } + lookupModel = account.GetMappedModel(modelID) + } else { + if !account.IsModelSupported(upstreamModel) { + continue + } + lookupModel = account.GetMappedModel(upstreamModel) + } + if strings.TrimSpace(lookupModel) != modelID { + publicAlias = true + } + metadata, ok := account.GetUpstreamModelMetadata(lookupModel) + if !ok { + if explicitTargetsConflict { + return codexModelMetadataOverride{ + reasoningConflict: true, + inputModalitiesConflict: true, + }, true + } + return codexModelMetadataOverride{}, false + } + candidates = append(candidates, metadata) + } + if len(candidates) == 0 { + return codexModelMetadataOverride{}, false + } + metadata := intersectUpstreamModelMetadata(modelID, candidates) + if publicAlias { + metadata.DisplayName = modelID + metadata.Description = configuredCodexCustomDescription + } + return metadata, true +} + +func codexExplicitModelTargetsConflict(accounts []Account, modelID string) bool { + targets := make(map[string]struct{}) + for i := range accounts { + account := &accounts[i] + mappedModel, matched := account.ResolveMappedModel(modelID) + mappedModel = strings.TrimSpace(mappedModel) + if !matched || mappedModel == "" { + continue + } + targets[strings.TrimSpace(account.Platform)+"\x00"+mappedModel] = struct{}{} + } + return len(targets) > 1 +} + +func codexExplicitModelTargetsConflictForPlatform(accounts []Account, platform, modelID string) bool { + targets := make(map[string]struct{}) + for i := range accounts { + account := &accounts[i] + if account.Platform != platform { + continue + } + mappedModel, matched := account.ResolveMappedModel(modelID) + mappedModel = strings.TrimSpace(mappedModel) + if !matched || mappedModel == "" { + continue + } + targets[mappedModel] = struct{}{} + } + return len(targets) > 1 +} + +func intersectUpstreamModelMetadata(modelID string, candidates []UpstreamModelMetadata) codexModelMetadataOverride { + result := codexModelMetadataOverride{UpstreamModelMetadata: UpstreamModelMetadata{ID: strings.TrimSpace(modelID)}} + for _, candidate := range candidates { + if result.DisplayName == "" && strings.TrimSpace(candidate.DisplayName) != "" { + result.DisplayName = strings.TrimSpace(candidate.DisplayName) + } + if result.Description == "" && strings.TrimSpace(candidate.Description) != "" { + result.Description = strings.TrimSpace(candidate.Description) + } + } + + reasoningKnown := true + reasoningValue := false + for i, candidate := range candidates { + if candidate.Reasoning == nil { + reasoningKnown = false + break + } + if i == 0 { + reasoningValue = *candidate.Reasoning + continue + } + if reasoningValue != *candidate.Reasoning { + reasoningKnown = false + result.reasoningConflict = true + break + } + } + if reasoningKnown { + result.Reasoning = &reasoningValue + if reasoningValue { + levels := normalizeReasoningLevels(candidates[0].SupportedReasoningLevels) + for _, candidate := range candidates[1:] { + levels = intersectOrderedStrings(levels, normalizeReasoningLevels(candidate.SupportedReasoningLevels)) + } + result.SupportedReasoningLevels = levels + if len(levels) == 0 { + result.reasoningConflict = true + } else { + sharedDefault := normalizeReasoningLevel(candidates[0].DefaultReasoningLevel) + for _, candidate := range candidates[1:] { + if normalizeReasoningLevel(candidate.DefaultReasoningLevel) != sharedDefault { + sharedDefault = "" + break + } + } + if !stringSliceContains(levels, sharedDefault) { + sharedDefault = levels[0] + } + result.DefaultReasoningLevel = sharedDefault + } + } + } + + modalitiesKnown := true + modalities := normalizeCodexInputModalities(candidates[0].InputModalities) + if len(modalities) == 0 { + modalitiesKnown = false + } + for _, candidate := range candidates[1:] { + candidateModalities := normalizeCodexInputModalities(candidate.InputModalities) + if len(candidateModalities) == 0 { + modalitiesKnown = false + break + } + modalities = intersectOrderedStrings(modalities, candidateModalities) + } + if modalitiesKnown && len(modalities) > 0 { + result.InputModalities = modalities + } else if modalitiesKnown { + result.inputModalitiesConflict = true + } + + contextKnown := true + for i, candidate := range candidates { + if candidate.ContextWindow <= 0 { + contextKnown = false + break + } + if i == 0 || candidate.ContextWindow < result.ContextWindow { + result.ContextWindow = candidate.ContextWindow + } + } + if !contextKnown { + result.ContextWindow = 0 + } + return result +} + +func applyUpstreamModelMetadataToCodexDescriptor( + descriptor *configuredCodexModelDescriptor, + metadata codexModelMetadataOverride, +) { + if descriptor == nil { + return + } + if strings.TrimSpace(metadata.DisplayName) != "" { + descriptor.DisplayName = strings.TrimSpace(metadata.DisplayName) + } + if strings.TrimSpace(metadata.Description) != "" { + descriptor.Description = strings.TrimSpace(metadata.Description) + } + if metadata.reasoningConflict { + descriptor.DefaultReasoningLevel = nil + descriptor.SupportedReasoningLevels = []configuredCodexReasoningLevel{} + } else if metadata.Reasoning != nil && !*metadata.Reasoning { + none := "none" + descriptor.DefaultReasoningLevel = &none + descriptor.SupportedReasoningLevels = []configuredCodexReasoningLevel{{ + Effort: "none", + Description: configuredCodexReasoningLevelDescription("none"), + }} + } else if metadata.Reasoning != nil && *metadata.Reasoning { + levels := normalizeReasoningLevels(metadata.SupportedReasoningLevels) + if len(levels) == 0 { + descriptor.DefaultReasoningLevel = nil + descriptor.SupportedReasoningLevels = []configuredCodexReasoningLevel{} + } else { + defaultLevel := normalizeReasoningLevel(metadata.DefaultReasoningLevel) + if !stringSliceContains(levels, defaultLevel) { + defaultLevel = levels[0] + } + descriptor.DefaultReasoningLevel = &defaultLevel + descriptor.SupportedReasoningLevels = make([]configuredCodexReasoningLevel, 0, len(levels)) + for _, level := range levels { + descriptor.SupportedReasoningLevels = append(descriptor.SupportedReasoningLevels, configuredCodexReasoningLevel{ + Effort: level, + Description: configuredCodexReasoningLevelDescription(level), + }) + } + } + } + if metadata.inputModalitiesConflict { + descriptor.InputModalities = []string{"text"} + } else if modalities := normalizeCodexInputModalities(metadata.InputModalities); len(modalities) > 0 { + descriptor.InputModalities = modalities + } + if metadata.ContextWindow > 0 { + descriptor.ContextWindow = metadata.ContextWindow + descriptor.MaxContextWindow = metadata.ContextWindow + } +} + +func configuredCodexReasoningLevelDescription(level string) string { + switch level { + case "none": + return "Use the model's default behavior without configurable reasoning" + case "minimal": + return "Minimal reasoning for the fastest responses" + case "low": + return "Fast responses with lighter reasoning" + case "medium": + return "Balanced reasoning for most coding tasks" + case "high": + return "Greater reasoning depth for coding and agent tasks" + case "xhigh": + return "Extra-high reasoning depth for difficult tasks" + case "max": + return "Maximum reasoning depth for complex tasks" + default: + return "Reasoning effort supported by the upstream model" + } +} + +func intersectOrderedStrings(left, right []string) []string { + rightSet := make(map[string]struct{}, len(right)) + for _, value := range right { + rightSet[value] = struct{}{} + } + intersection := make([]string, 0, len(left)) + for _, value := range left { + if _, ok := rightSet[value]; ok { + intersection = append(intersection, value) + } + } + return intersection +} + +func stringSliceContains(values []string, target string) bool { + if target == "" { + return false + } + for _, value := range values { + if value == target { + return true + } + } + return false +} diff --git a/backend/internal/service/openai_codex_model_metadata_test.go b/backend/internal/service/openai_codex_model_metadata_test.go new file mode 100644 index 0000000000..5161617d7b --- /dev/null +++ b/backend/internal/service/openai_codex_model_metadata_test.go @@ -0,0 +1,385 @@ +package service + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +// Scenario: mixed groups prefer capability metadata synced for the routed account. +func TestBuildCodexModelsManifestForGroupUsesSyncedAccountMetadata(t *testing.T) { + t.Parallel() + + const groupID int64 = 735 + account := Account{ + ID: 25, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://opencode.ai/zen/v1", + "model_mapping": map[string]any{"x-preview-f-free": "x-preview-f-free"}, + }, + Extra: map[string]any{ + UpstreamModelMetadataExtraKey: map[string]any{ + "source": "models.dev", + "models": map[string]any{ + "x-preview-f-free": map[string]any{ + "id": "x-preview-f-free", + "display_name": "Ox Alpha Free (Unlimited)", + "description": "Stealth reasoning model", + "reasoning": true, + "supported_reasoning_levels": []any{"low", "high", "max"}, + "input_modalities": []any{"text", "image"}, + "context_window": float64(1_000_000), + "max_output_tokens": float64(131_072), + }, + }, + }, + }, + } + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {account}, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"x-preview-f-free"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "Ox Alpha Free (Unlimited)", models[0]["display_name"]) + require.Equal(t, "low", models[0]["default_reasoning_level"]) + require.Equal(t, []string{"low", "high", "max"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) +} + +// Scenario: an explicitly non-reasoning model remains directly selectable in Codex. +func TestBuildCodexModelsManifestForGroupUsesNoneForExplicitNonReasoningMetadata(t *testing.T) { + t.Parallel() + + const groupID int64 = 737 + reasoning := false + account := Account{ + ID: 28, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{"company-coding-model": "company-coding-model"}, + }, + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "company-coding-model": { + ID: "company-coding-model", Reasoning: &reasoning, + InputModalities: []string{"text"}, ContextWindow: 64_000, + }, + }}) + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {account}, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformComposite}, "", []string{"company-coding-model"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "none", models[0]["default_reasoning_level"]) + require.Equal(t, []string{"none"}, effortsFromManifestModel(t, models[0])) +} + +// Scenario: multiple schedulable accounts advertise only their shared capabilities. +func TestBuildCodexModelsManifestForGroupIntersectsSyncedAccountMetadata(t *testing.T) { + t.Parallel() + + const groupID int64 = 736 + reasoning := true + newAccount := func(id int64, levels, modalities []string, contextWindow int64) Account { + account := Account{ + ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{"shared-model": "shared-model"}, + }, + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "shared-model": { + ID: "shared-model", Reasoning: &reasoning, + SupportedReasoningLevels: levels, + InputModalities: modalities, + ContextWindow: contextWindow, + }, + }}) + return account + } + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: { + newAccount(26, []string{"low", "high"}, []string{"text", "image"}, 256_000), + newAccount(27, []string{"high", "max"}, []string{"text"}, 128_000), + }, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformComposite}, "", []string{"shared-model"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []string{"high"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, "high", models[0]["default_reasoning_level"]) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.EqualValues(t, 128_000, models[0]["context_window"]) +} + +// Scenario: the same public alias may target different models on one platform when complete snapshots can be intersected. +func TestBuildCodexModelsManifestForGroupIntersectsDifferentMappedTargetsWithoutLeakingAlias(t *testing.T) { + t.Parallel() + + const groupID int64 = 739 + reasoning := true + newAccount := func(id int64, target, displayName, description string, levels, modalities []string, contextWindow int64) Account { + account := Account{ + ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{"my-coder": target}, + }, + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + target: { + ID: target, DisplayName: displayName, Description: description, Reasoning: &reasoning, + SupportedReasoningLevels: levels, + InputModalities: modalities, + ContextWindow: contextWindow, + }, + }}) + return account + } + openAIAccount := newAccount( + 31, + "gpt-5.6-sol", + "GPT-5.6 Sol", + "OpenAI upstream model", + []string{"low", "medium", "high", "xhigh"}, + []string{"text", "image"}, + 272_000, + ) + arkAccount := newAccount( + 32, + "glm-5.3", + "GLM 5.3", + "Ark upstream model", + []string{"low", "medium", "high"}, + []string{"text"}, + 1_000_000, + ) + + for _, accounts := range [][]Account{{openAIAccount, arkAccount}, {arkAccount, openAIAccount}} { + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: accounts, + }}} + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformOpenAI}, "", []string{"my-coder"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "my-coder", models[0]["slug"]) + require.Equal(t, "my-coder", models[0]["display_name"]) + require.Equal(t, "Custom model routed through Sub2API.", models[0]["description"]) + require.Equal(t, []string{"low", "medium", "high"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.EqualValues(t, 272_000, models[0]["context_window"]) + } +} + +// Scenario: temporarily unschedulable mapped accounts still participate in capability intersection. +func TestBuildCodexModelsManifestForGroupIntersectsUnschedulableMappedAccounts(t *testing.T) { + t.Parallel() + + const groupID int64 = 741 + schedulable := newCodexCatalogMappedAccount( + 41, + "gpt-5.6-sol", + "GPT-5.6 Sol", + []string{"low", "medium", "high", "xhigh"}, + []string{"text", "image"}, + 1_000_000, + true, + nil, + ) + unschedulable := newCodexCatalogMappedAccount( + 42, + "glm-5.3", + "GLM 5.3", + []string{"low", "medium", "high"}, + []string{"text"}, + 272_000, + false, + map[string]any{"exclusive-model": "exclusive-upstream"}, + ) + svc := &GatewayService{accountRepo: splitCodexModelsAccountRepo{ + schedulable: map[int64][]Account{groupID: {schedulable}}, + catalog: map[int64][]Account{groupID: {schedulable, unschedulable}}, + }} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformOpenAI}, "", []string{"my-coder"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "my-coder", models[0]["slug"]) + require.Equal(t, "my-coder", models[0]["display_name"]) + require.Equal(t, []string{"low", "medium", "high"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.EqualValues(t, 272_000, models[0]["context_window"]) +} + +// Scenario: deleting an account can widen the advertised contract. +func TestBuildCodexModelsManifestForGroupWidensAfterUnschedulableAccountIsRemoved(t *testing.T) { + t.Parallel() + + const groupID int64 = 742 + remaining := newCodexCatalogMappedAccount( + 41, + "gpt-5.6-sol", + "GPT-5.6 Sol", + []string{"low", "medium", "high", "xhigh"}, + []string{"text", "image"}, + 1_000_000, + true, + nil, + ) + svc := &GatewayService{accountRepo: splitCodexModelsAccountRepo{ + schedulable: map[int64][]Account{groupID: {remaining}}, + catalog: map[int64][]Account{groupID: {remaining}}, + }} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformOpenAI}, "", []string{"my-coder"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) +} + +func TestBuildCodexModelsManifestForGroupFallsBackToSchedulableWhenListByGroupFails(t *testing.T) { + t.Parallel() + + const groupID int64 = 743 + repo := &countingCodexModelsAccountRepo{ + accounts: []Account{newCodexCatalogMappedAccount( + 41, + "gpt-5.6-sol", + "GPT-5.6 Sol", + []string{"low", "medium", "high", "xhigh"}, + []string{"text", "image"}, + 1_000_000, + true, + nil, + )}, + listByGroupErr: errors.New("group listing unavailable"), + } + svc := &GatewayService{accountRepo: repo} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformOpenAI}, "", []string{"my-coder"}, + ) + require.NoError(t, err) + require.Equal(t, int32(1), repo.calls.Load()) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) +} + +// Scenario: a Composite alias claimed across platforms remains ambiguous and fails closed. +func TestBuildCodexModelsManifestForGroupKeepsCrossPlatformAliasAmbiguityClosed(t *testing.T) { + t.Parallel() + + const groupID int64 = 740 + reasoning := true + newAccount := func(id int64, platform, target string) Account { + account := Account{ + ID: id, Platform: platform, Type: AccountTypeAPIKey, + Credentials: map[string]any{"model_mapping": map[string]any{"shared-alias": target}}, + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + target: { + ID: target, DisplayName: target, Reasoning: &reasoning, + SupportedReasoningLevels: []string{"low", "high"}, + InputModalities: []string{"text", "image"}, + ContextWindow: 128_000, + }, + }}) + return account + } + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: { + newAccount(33, PlatformOpenAI, "gpt-5.6-sol"), + newAccount(34, PlatformGrok, "grok-4.6"), + }, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformComposite}, "", []string{"shared-alias"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "shared-alias", models[0]["display_name"]) + require.Empty(t, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) +} + +func TestBuildCodexModelsManifestForGroupDoesNotAdvertiseNoneWhenAccountReasoningConflicts(t *testing.T) { + t.Parallel() + + const groupID int64 = 738 + reasoning := true + noReasoning := false + newAccount := func(id int64, metadata UpstreamModelMetadata) Account { + account := Account{ + ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{"shared-model": "shared-model"}, + }, + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "shared-model": metadata, + }}) + return account + } + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: { + newAccount(29, UpstreamModelMetadata{ + ID: "shared-model", Reasoning: &reasoning, + SupportedReasoningLevels: []string{"low", "high"}, + InputModalities: []string{"text"}, ContextWindow: 128_000, + }), + newAccount(30, UpstreamModelMetadata{ + ID: "shared-model", Reasoning: &noReasoning, + InputModalities: []string{"text"}, ContextWindow: 128_000, + }), + }, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), &Group{ID: groupID, Platform: PlatformComposite}, "", []string{"shared-model"}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + _, hasDefault := models[0]["default_reasoning_level"] + require.False(t, hasDefault) + require.Empty(t, models[0]["supported_reasoning_levels"]) +} diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index 7c06dc99bd..95be9ebbaf 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -16,8 +16,11 @@ import ( "sync" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "golang.org/x/net/http2" "golang.org/x/sync/singleflight" ) @@ -32,14 +35,1191 @@ const ( codexModelsManifestCacheTTL = 30 * time.Second codexModelsManifestCacheStaleTTL = 5 * time.Minute codexModelsManifestRequestTimeout = 15 * time.Second + codexAutoModelPrefix = "codex-auto-" ) +// FilterCodexModelIDsForGroup removes dedicated media-generation models, +// wildcard mapping keys, and Codex automatic modes from a client catalog. +// Automatic modes are retained only when the group's enabled custom model list +// explicitly selects the exact slug; account model mappings describe routing +// and are not feature opt-ins. Wildcard keys such as "foo-*" are routing +// patterns, not concrete Codex models. +func FilterCodexModelIDsForGroup(modelIDs []string, group *Group) []string { + explicitlyEnabled := make(map[string]struct{}) + if group != nil && group.CustomModelsListEnabled() { + for _, modelID := range group.ModelsListConfig.Models { + modelID = strings.TrimSpace(modelID) + if strings.HasPrefix(modelID, codexAutoModelPrefix) { + explicitlyEnabled[modelID] = struct{}{} + } + } + } + + filtered := make([]string, 0, len(modelIDs)) + for _, modelID := range modelIDs { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + continue + } + if isCodexDedicatedMediaModel(modelID) { + continue + } + if strings.Contains(modelID, "*") { + continue + } + if strings.HasPrefix(modelID, codexAutoModelPrefix) { + if _, ok := explicitlyEnabled[modelID]; !ok { + continue + } + } + filtered = append(filtered, modelID) + } + return filtered +} + +func isCodexDedicatedMediaModel(modelID string) bool { + canonical := codexProviderQualifiedModelID(modelID) + return IsGPTImageGenerationModel(canonical) || + isImageGenerationModel(canonical) || + xai.IsGrokImagineModel(modelID) +} + +func codexProviderQualifiedModelID(modelID string) string { + modelID = strings.TrimSpace(modelID) + if slash := strings.LastIndexByte(modelID, '/'); slash >= 0 { + modelID = strings.TrimSpace(modelID[slash+1:]) + } + return strings.TrimPrefix(modelID, "models/") +} + // CodexModelsManifest carries the client representation plus caching metadata. type CodexModelsManifest struct { - Body []byte - ETag string - upstreamETag string - NotModified bool + Body []byte + ETag string + upstreamETag string + upstreamSourceBody []byte + convertedFromOpenAIModelList bool + NotModified bool +} + +// BuildGroupConfiguredCodexModelsManifest builds a Codex catalog exclusively +// from the public model names configured on accounts in an OpenAI group. The +// boolean result distinguishes "no explicit configuration" from a configured +// catalog that becomes empty after group-level filtering. +func (s *OpenAIGatewayService) BuildGroupConfiguredCodexModelsManifest( + ctx context.Context, + group *Group, + ifNoneMatch string, +) (*CodexModelsManifest, bool, error) { + if s == nil || s.accountRepo == nil || group == nil || group.Platform != PlatformOpenAI { + return nil, false, nil + } + + visible, catalog, err := loadCodexGroupCatalogAccounts(ctx, s.accountRepo, group.ID) + if err != nil { + return nil, false, fmt.Errorf("load group configured Codex models: %w", err) + } + configuredModels := openAIConfiguredCodexModelIDsForGroup(visible, group) + if len(configuredModels) == 0 { + return nil, false, nil + } + + body, err := buildCodexModelsManifestForAccounts( + PlatformOpenAI, + configuredModels, + catalog, + nil, + true, + ) + if err != nil { + return nil, false, fmt.Errorf("initialize group configured Codex models: %w", err) + } + body, _, err = mergeConfiguredCodexModelsManifest( + body, + nil, + group.ModelsListConfig.Models, + group.CustomModelsListEnabled(), + ) + if err != nil { + return nil, false, fmt.Errorf("build group configured Codex models: %w", err) + } + manifest := &CodexModelsManifest{ + Body: body, + ETag: codexModelsManifestBodyETag(body), + } + if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { + manifest.Body = nil + manifest.NotModified = true + } + return manifest, true, nil +} + +// MergeGroupConfiguredCodexModels adds account model aliases that are visible +// to the authenticated OpenAI group without discarding metadata from upstream +// Codex model entries. A group's custom models list also filters the picker, +// matching the standard /v1/models display policy. +func (s *OpenAIGatewayService) MergeGroupConfiguredCodexModels( + ctx context.Context, + group *Group, + manifest *CodexModelsManifest, + ifNoneMatch string, +) error { + if s == nil || s.accountRepo == nil || group == nil || manifest == nil || manifest.NotModified { + return nil + } + if group.Platform != PlatformOpenAI || len(manifest.Body) == 0 { + return nil + } + + configuredModels, err := s.groupConfiguredCodexModelIDs(ctx, group) + if err != nil { + return fmt.Errorf("load group configured Codex models: %w", err) + } + body, changed, err := mergeConfiguredCodexModelsManifest( + manifest.Body, + configuredModels, + group.ModelsListConfig.Models, + group.CustomModelsListEnabled(), + ) + if err != nil { + return fmt.Errorf("merge group configured Codex models: %w", err) + } + if changed { + manifest.Body = body + manifest.ETag = codexModelsManifestBodyETag(body) + } + if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { + manifest.Body = nil + manifest.NotModified = true + } + return nil +} + +func (s *OpenAIGatewayService) groupConfiguredCodexModelIDs(ctx context.Context, group *Group) ([]string, error) { + if group == nil { + return nil, nil + } + accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, group.ID) + if err != nil { + return nil, err + } + return openAIConfiguredCodexModelIDsForGroup(accounts, group), nil +} + +// loadCodexGroupCatalogAccounts separates picker membership from capability +// intersection. visible accounts are currently schedulable and decide which +// public aliases appear. catalog accounts are all non-deleted active group +// members; their snapshots keep advertised capabilities from widening when a +// mapped account is only temporarily unschedulable. If ListByGroup fails, the +// catalog falls back to the schedulable set so a listing error does not fail +// the client request. +func loadCodexGroupCatalogAccounts(ctx context.Context, repo AccountRepository, groupID int64) (visible []Account, catalog []Account, err error) { + if repo == nil { + return nil, nil, nil + } + visible, err = repo.ListSchedulableByGroupID(ctx, groupID) + if err != nil { + return nil, nil, err + } + catalog = visible + groupAccounts, listErr := repo.ListByGroup(ctx, groupID) + if listErr != nil { + return visible, catalog, nil + } + return visible, groupAccounts, nil +} + +func openAIConfiguredCodexModelIDs(accounts []Account) []string { + seen := make(map[string]struct{}) + models := make([]string, 0) + for i := range accounts { + account := &accounts[i] + if account.Platform != PlatformOpenAI { + continue + } + for modelID := range account.GetModelMapping() { + modelID = strings.TrimSpace(modelID) + if modelID == "" || strings.Contains(modelID, "*") { + continue + } + if _, exists := seen[modelID]; exists { + continue + } + seen[modelID] = struct{}{} + models = append(models, modelID) + } + } + sort.Strings(models) + return models +} + +func openAIConfiguredCodexModelIDsForGroup(accounts []Account, group *Group) []string { + models := openAIConfiguredCodexModelIDs(accounts) + if group == nil || !group.CustomModelsListEnabled() { + return models + } + + seen := make(map[string]struct{}, len(models)+len(group.ModelsListConfig.Models)) + for _, modelID := range models { + seen[modelID] = struct{}{} + } + for _, selectedModel := range group.ModelsListConfig.Models { + selectedModel = strings.TrimSpace(selectedModel) + if selectedModel == "" || strings.Contains(selectedModel, "*") { + continue + } + for i := range accounts { + account := &accounts[i] + if account.Platform != PlatformOpenAI { + continue + } + mappedModel, matched := account.ResolveMappedModel(selectedModel) + if !matched || strings.TrimSpace(mappedModel) == "" { + continue + } + if _, exists := seen[selectedModel]; !exists { + seen[selectedModel] = struct{}{} + models = append(models, selectedModel) + } + break + } + } + sort.Strings(models) + return models +} + +const ( + configuredCodexModelPriority = 50 + configuredCodexCustomDescription = "Custom model routed through Sub2API." + configuredCodexFallbackContext = 272_000 + configuredCodexDeepSeekV4Context = 1_000_000 + configuredCodexGrokContext = 500_000 + configuredCodexGrokBuildContext = 256_000 + configuredCodexGPT56MaxContext = 872_000 + configuredCodexToolOutputMaxTokens = 10_000 +) + +type configuredCodexReasoningLevel struct { + Effort string `json:"effort"` + Description string `json:"description"` +} + +type configuredCodexTruncationPolicy struct { + Mode string `json:"mode"` + Limit int64 `json:"limit"` +} + +type configuredCodexModelMessages struct { + InstructionsTemplate string `json:"instructions_template"` + InstructionsVariables any `json:"instructions_variables"` + Approvals any `json:"approvals"` + CollaborationModes any `json:"collaboration_modes"` + AutoReview any `json:"auto_review"` + Permissions any `json:"permissions"` + MultiAgent any `json:"multi_agent"` + TokenBudget any `json:"token_budget"` + GuardianV2 any `json:"guardian_v2"` +} + +// configuredCodexModelDescriptor is the minimum complete ModelInfo contract +// understood by current Codex clients. Several nullable fields are intentionally +// emitted: unlike ordinary OpenAI /v1/models entries, the Codex manifest parser +// requires them to be present. +type configuredCodexModelDescriptor struct { + Slug string `json:"slug"` + DisplayName string `json:"display_name"` + Description string `json:"description"` + DefaultReasoningLevel *string `json:"default_reasoning_level,omitempty"` + SupportedReasoningLevels []configuredCodexReasoningLevel `json:"supported_reasoning_levels"` + ShellType string `json:"shell_type"` + Visibility string `json:"visibility"` + SupportedInAPI bool `json:"supported_in_api"` + Priority int `json:"priority"` + AdditionalSpeedTiers []string `json:"additional_speed_tiers"` + ServiceTiers []any `json:"service_tiers"` + DefaultServiceTier any `json:"default_service_tier"` + AvailabilityNUX any `json:"availability_nux"` + Upgrade any `json:"upgrade"` + ModelMessages configuredCodexModelMessages `json:"model_messages"` + IncludeSkillsUsageInstructions bool `json:"include_skills_usage_instructions"` + IncludePluginUsageInstructions bool `json:"include_plugin_usage_instructions"` + IncludeAppsUsageInstructions bool `json:"include_apps_usage_instructions"` + SupportsReasoningSummaryParameter bool `json:"supports_reasoning_summary_parameter"` + DefaultReasoningSummary string `json:"default_reasoning_summary"` + SupportVerbosity bool `json:"support_verbosity"` + DefaultVerbosity *string `json:"default_verbosity"` + ApplyPatchToolType *string `json:"apply_patch_tool_type"` + WebSearchToolType string `json:"web_search_tool_type"` + TruncationPolicy configuredCodexTruncationPolicy `json:"truncation_policy"` + SupportsImageDetailOriginal bool `json:"supports_image_detail_original"` + SupportsParallelToolCalls bool `json:"supports_parallel_tool_calls"` + ContextWindow int64 `json:"context_window"` + MaxContextWindow int64 `json:"max_context_window"` + AutoCompactTokenLimit any `json:"auto_compact_token_limit"` + CompHash any `json:"comp_hash"` + EffectiveContextWindowPercent int64 `json:"effective_context_window_percent"` + ExperimentalSupportedTools []string `json:"experimental_supported_tools"` + InputModalities []string `json:"input_modalities"` + SupportsSearchTool bool `json:"supports_search_tool"` + UseResponsesLite bool `json:"use_responses_lite"` + NodeREPLAutoReviewRequired bool `json:"node_repl_auto_review_required"` + NodeREPLDisabled bool `json:"node_repl_disabled"` + AutoReviewModelOverride any `json:"auto_review_model_override"` + ModelSpecialty any `json:"model_specialty"` + ToolMode any `json:"tool_mode"` + MultiAgentVersion any `json:"multi_agent_version"` +} + +type codexModelMetadataOverride struct { + UpstreamModelMetadata + reasoningConflict bool + inputModalitiesConflict bool +} + +func newConfiguredCodexModelDescriptor(modelID string) configuredCodexModelDescriptor { + modelID = strings.TrimSpace(modelID) + noReasoningLevel := "none" + descriptor := configuredCodexModelDescriptor{ + Slug: modelID, + DisplayName: modelID, + Description: configuredCodexCustomDescription, + DefaultReasoningLevel: &noReasoningLevel, + SupportedReasoningLevels: []configuredCodexReasoningLevel{ + {Effort: "none", Description: configuredCodexReasoningLevelDescription("none")}, + }, + ShellType: "unified_exec", + Visibility: "list", + SupportedInAPI: true, + Priority: configuredCodexModelPriority, + AdditionalSpeedTiers: []string{}, + ServiceTiers: []any{}, + ModelMessages: configuredCodexModelMessages{InstructionsTemplate: openai.CodexBaseInstructionsForModel(modelID)}, + SupportsReasoningSummaryParameter: true, + DefaultReasoningSummary: "auto", + WebSearchToolType: "text", + TruncationPolicy: configuredCodexTruncationPolicy{Mode: "bytes", Limit: configuredCodexToolOutputMaxTokens}, + ContextWindow: configuredCodexFallbackContext, + MaxContextWindow: configuredCodexFallbackContext, + EffectiveContextWindowPercent: 95, + ExperimentalSupportedTools: []string{}, + InputModalities: []string{"text"}, + } + + if isDeepSeekCodexModel(modelID) { + defaultReasoningLevel := "high" + descriptor.DisplayName = deepSeekCodexDisplayName(modelID) + descriptor.Description = "DeepSeek coding and reasoning model routed through Sub2API." + descriptor.DefaultReasoningLevel = &defaultReasoningLevel + descriptor.SupportedReasoningLevels = []configuredCodexReasoningLevel{ + {Effort: "low", Description: "Fast responses with lighter reasoning"}, + {Effort: "high", Description: "Greater reasoning depth for coding and agent tasks"}, + {Effort: "max", Description: "Maximum reasoning depth for complex tasks"}, + } + descriptor.SupportsParallelToolCalls = true + descriptor.ContextWindow = configuredCodexDeepSeekV4Context + descriptor.MaxContextWindow = configuredCodexDeepSeekV4Context + } + + if isGrokCodexModel(modelID) { + descriptor.DisplayName = grokCodexDisplayName(modelID) + descriptor.Description = "Grok coding and reasoning model routed through Sub2API." + descriptor.SupportsParallelToolCalls = true + descriptor.ContextWindow = grokCodexContextWindow(modelID) + descriptor.MaxContextWindow = descriptor.ContextWindow + if grokCodexSupportsReasoningEffort(modelID) { + defaultReasoningLevel := "high" + descriptor.DefaultReasoningLevel = &defaultReasoningLevel + descriptor.SupportedReasoningLevels = configuredCodexGrokReasoningLevels(modelID) + } + } + + if isClaudeCodexModel(modelID) { + descriptor.DisplayName = claudeCodexDisplayName(modelID) + descriptor.Description = "Claude coding and reasoning model routed through Sub2API." + descriptor.SupportsParallelToolCalls = true + if levels := configuredCodexClaudeReasoningLevels(modelID); len(levels) > 0 { + defaultReasoningLevel := claudeCodexDefaultReasoningLevel(levels) + descriptor.DefaultReasoningLevel = &defaultReasoningLevel + descriptor.SupportedReasoningLevels = levels + } + } + + if isOpenAICodexGPTModel(modelID) { + descriptor.DisplayName = openaiCodexDisplayName(modelID) + descriptor.Description = "OpenAI GPT coding model routed through Sub2API." + descriptor.SupportsParallelToolCalls = true + if isOpenAICodexReasoningGPTModel(modelID) { + defaultReasoningLevel := "medium" + if getNormalizedCodexModel(modelID) == "gpt-5.6-sol" { + defaultReasoningLevel = "low" + } + descriptor.DefaultReasoningLevel = &defaultReasoningLevel + descriptor.SupportedReasoningLevels = configuredCodexGPTReasoningLevels(modelID) + descriptor.DefaultReasoningSummary = "none" + descriptor.TruncationPolicy = configuredCodexTruncationPolicy{Mode: "tokens", Limit: configuredCodexToolOutputMaxTokens} + if isOpenAIGPT56Model(modelID) { + descriptor.MaxContextWindow = configuredCodexGPT56MaxContext + } + } + if SupportsVerbosity(modelID) { + defaultVerbosity := "low" + descriptor.SupportVerbosity = true + descriptor.DefaultVerbosity = &defaultVerbosity + } + } + + return descriptor +} + +func configuredCodexGrokReasoningLevels(modelID string) []configuredCodexReasoningLevel { + levels := []configuredCodexReasoningLevel{ + {Effort: "low", Description: "Fast responses with lighter reasoning"}, + {Effort: "medium", Description: "Balanced reasoning for most coding tasks"}, + {Effort: "high", Description: "Greater reasoning depth for coding and agent tasks"}, + } + if grokSupportsXHighReasoningEffort(modelID) { + levels = append(levels, configuredCodexReasoningLevel{ + Effort: "xhigh", + Description: "Extra-high reasoning depth for difficult tasks", + }) + } + return levels +} + +func configuredCodexClaudeReasoningLevels(modelID string) []configuredCodexReasoningLevel { + descriptions := map[string]string{ + "low": "Fast responses with lighter reasoning", + "medium": "Balanced reasoning for most coding tasks", + "high": "Greater reasoning depth for coding and agent tasks", + "xhigh": "Extra-high reasoning depth for difficult tasks", + "max": "Maximum reasoning depth for complex tasks", + } + levels := claude.EffortLevelsForModel(modelID) + out := make([]configuredCodexReasoningLevel, 0, len(levels)) + for _, effort := range levels { + out = append(out, configuredCodexReasoningLevel{ + Effort: effort, + Description: descriptions[effort], + }) + } + return out +} + +func claudeCodexDefaultReasoningLevel(levels []configuredCodexReasoningLevel) string { + for _, preferred := range []string{"medium", "high", "low"} { + for _, level := range levels { + if level.Effort == preferred { + return preferred + } + } + } + if len(levels) == 0 { + return "" + } + return levels[0].Effort +} + +func configuredCodexGPTReasoningLevels(modelID string) []configuredCodexReasoningLevel { + levels := []configuredCodexReasoningLevel{ + {Effort: "low", Description: "Fast responses with lighter reasoning"}, + {Effort: "medium", Description: "Balanced reasoning for most coding tasks"}, + {Effort: "high", Description: "Greater reasoning depth for coding and agent tasks"}, + {Effort: "xhigh", Description: "Extra-high reasoning depth for difficult tasks"}, + } + normalized := getNormalizedCodexModel(modelID) + if isOpenAIGPT56Model(modelID) { + levels = append(levels, configuredCodexReasoningLevel{ + Effort: "max", + Description: "Maximum reasoning depth for complex tasks", + }) + } + if normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" { + levels = append(levels, configuredCodexReasoningLevel{ + Effort: "ultra", + Description: "Maximum reasoning with automatic task delegation", + }) + } + return levels +} + +func isOpenAICodexGPTModel(modelID string) bool { + normalized := canonicalizeOpenAIModelAliasSpelling(modelID) + if normalized == "" || strings.HasPrefix(normalized, "gpt-image") { + return false + } + return strings.HasPrefix(normalized, "gpt-") +} + +func isOpenAICodexReasoningGPTModel(modelID string) bool { + normalized := canonicalizeOpenAIModelAliasSpelling(modelID) + return strings.HasPrefix(normalized, "gpt-5") +} + +func isOpenAICodexImageInputModel(modelID string) bool { + normalized := canonicalizeOpenAIModelAliasSpelling(modelID) + return strings.HasPrefix(normalized, "gpt-5") || + strings.HasPrefix(normalized, "gpt-4o") || + strings.HasPrefix(normalized, "gpt-4.1") || + strings.HasPrefix(normalized, "gpt-4.5") || + strings.HasPrefix(normalized, "gpt-4-turbo") || + strings.HasPrefix(normalized, "gpt-4-vision") +} + +func isOfficialOpenAICodexCatalogModel(modelID string) bool { + normalized := strings.ToLower(codexProviderQualifiedModelID(modelID)) + if normalized == "" || isCodexDedicatedMediaModel(normalized) { + return false + } + if strings.HasPrefix(normalized, "codex-") { + return true + } + if strings.HasPrefix(normalized, "o1") || strings.HasPrefix(normalized, "o3") || strings.HasPrefix(normalized, "o4") { + return true + } + if !strings.HasPrefix(normalized, "gpt-") { + return false + } + for _, incompatibleFamily := range []string{"audio", "realtime", "transcribe", "tts"} { + if strings.Contains(normalized, incompatibleFamily) { + return false + } + } + return true +} + +func openaiCodexDisplayName(modelID string) string { + normalized := canonicalizeOpenAIModelAliasSpelling(modelID) + if normalized == "" { + return modelID + } + for _, model := range openai.DefaultModels { + if strings.EqualFold(model.ID, normalized) && strings.TrimSpace(model.DisplayName) != "" { + return model.DisplayName + } + } + return modelID +} + +func deepSeekCodexDisplayName(modelID string) string { + switch strings.ToLower(strings.TrimSpace(modelID)) { + case "deepseek-v4-pro", "deepseek-4-pro": + return "DeepSeek V4 Pro" + case "deepseek-v4-flash", "deepseek-4-flash": + return "DeepSeek V4 Flash" + default: + return modelID + } +} + +func isDeepSeekCodexModel(modelID string) bool { + return strings.HasPrefix(strings.ToLower(strings.TrimSpace(modelID)), "deepseek-") +} + +func isGrokCodexModel(modelID string) bool { + return xai.IsGrokModelID(modelID) +} + +func grokCodexSupportsReasoningEffort(modelID string) bool { + if grokSupportsReasoningEffort(modelID) { + return true + } + canonical := xai.ResolveGrokTextResponsesModelID(modelID) + if canonical == "" || strings.EqualFold(canonical, modelID) { + return false + } + return grokSupportsReasoningEffort(canonical) +} + +func grokCodexDisplayName(modelID string) string { + normalized := strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(modelID))) + if normalized == "" { + return modelID + } + if name := grokDefaultDisplayName(normalized); name != "" { + return name + } + canonical := strings.ToLower(xai.ResolveGrokTextResponsesModelID(normalized)) + if canonical != "" && canonical != normalized { + if name := grokDefaultDisplayName(canonical); name != "" { + return name + } + } + return modelID +} + +func grokDefaultDisplayName(modelID string) string { + for _, model := range xai.DefaultModels() { + if model.ID == modelID { + return strings.TrimSpace(model.DisplayName) + } + } + return "" +} + +func grokCodexContextWindow(modelID string) int64 { + normalized := strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(modelID))) + if strings.HasPrefix(normalized, "grok-build") { + return configuredCodexGrokBuildContext + } + return configuredCodexGrokContext +} + +func isClaudeCodexModel(modelID string) bool { + platform, detected := DetectModelPlatform(modelID) + return detected && platform == PlatformAnthropic +} + +func claudeCodexDisplayName(modelID string) string { + normalized := strings.ToLower(codexProviderQualifiedModelID(modelID)) + normalized = strings.TrimPrefix(normalized, "anthropic.") + if normalized == "" { + return modelID + } + for _, model := range claude.DefaultModels { + if strings.EqualFold(model.ID, normalized) && strings.TrimSpace(model.DisplayName) != "" { + return model.DisplayName + } + } + if canonical, ok := claude.ModelIDOverrides[normalized]; ok { + for _, model := range claude.DefaultModels { + if model.ID == canonical && strings.TrimSpace(model.DisplayName) != "" { + return model.DisplayName + } + } + } + return modelID +} + +// BuildCodexModelsManifest builds a standalone Codex model catalog for models +// routed through a custom provider. The response is also suitable for saving +// as model_catalog_json in clients that do not refresh custom-provider catalogs. +func BuildCodexModelsManifest(modelIDs []string) ([]byte, error) { + return buildCodexModelsManifest(modelIDs, nil, nil, nil) +} + +// BuildCodexModelsManifestForGroup derives input capabilities from the +// concrete Responses route and group accounts behind a group. Unknown or mixed +// capabilities fail closed to the text-only descriptor used by the standalone +// builder. Caller-supplied model IDs still decide which slugs appear; advertised +// capabilities intersect all active group members that map the alias, including +// accounts that are not currently schedulable. +func (s *GatewayService) BuildCodexModelsManifestForGroup( + ctx context.Context, + group *Group, + platformOverride string, + modelIDs []string, +) ([]byte, error) { + if s == nil || s.accountRepo == nil || group == nil { + return BuildCodexModelsManifest(modelIDs) + } + effectivePlatform := strings.TrimSpace(platformOverride) + if effectivePlatform == "" { + effectivePlatform = group.Platform + } + if effectivePlatform != PlatformComposite && !isConcreteRequestPlatform(effectivePlatform) { + return BuildCodexModelsManifest(modelIDs) + } + + _, catalog, err := loadCodexGroupCatalogAccounts(ctx, s.accountRepo, group.ID) + if err != nil { + return BuildCodexModelsManifest(modelIDs) + } + var compositeRoutes []CompositeModelRoute + compositeRoutesAvailable := true + if effectivePlatform == PlatformComposite && s.compositeResolver != nil && s.compositeResolver.repo != nil { + compositeRoutes, err = s.compositeResolver.repo.ListByGroup(ctx, group.ID, false) + if err != nil { + compositeRoutesAvailable = false + } + } + return buildCodexModelsManifestForAccounts( + effectivePlatform, + modelIDs, + catalog, + compositeRoutes, + compositeRoutesAvailable, + ) +} + +func buildCodexModelsManifestForAccounts( + effectivePlatform string, + modelIDs []string, + accounts []Account, + compositeRoutes []CompositeModelRoute, + compositeRoutesAvailable bool, +) ([]byte, error) { + imageInputModels := make(map[string]bool, len(modelIDs)) + metadataModels := codexCatalogMetadataModels( + effectivePlatform, + modelIDs, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ) + modelMetadata := make(map[string]codexModelMetadataOverride, len(modelIDs)) + for _, modelID := range modelIDs { + modelID = strings.TrimSpace(modelID) + if groupCodexModelSupportsImageInput( + effectivePlatform, + modelID, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ) { + imageInputModels[modelID] = true + } + if metadata, ok := groupCodexModelMetadata( + effectivePlatform, + modelID, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ); ok { + modelMetadata[modelID] = metadata + } + } + return buildCodexModelsManifest(modelIDs, imageInputModels, metadataModels, modelMetadata) +} + +func buildCodexModelsManifest( + modelIDs []string, + imageInputModels map[string]bool, + metadataModels map[string]string, + modelMetadata map[string]codexModelMetadataOverride, +) ([]byte, error) { + seen := make(map[string]struct{}, len(modelIDs)) + models := make([]configuredCodexModelDescriptor, 0, len(modelIDs)) + for _, modelID := range modelIDs { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + continue + } + if _, exists := seen[modelID]; exists { + continue + } + metadataModelID := strings.TrimSpace(metadataModels[modelID]) + if metadataModelID == "" { + metadataModelID = modelID + } + if isCodexDedicatedMediaModel(modelID) || isCodexDedicatedMediaModel(metadataModelID) { + continue + } + seen[modelID] = struct{}{} + descriptor := newConfiguredCodexModelDescriptor(metadataModelID) + descriptor.Slug = modelID + if imageInputModels[modelID] { + descriptor.InputModalities = []string{"text", "image"} + } + if metadata, ok := modelMetadata[modelID]; ok { + applyUpstreamModelMetadataToCodexDescriptor(&descriptor, metadata) + } + if metadataModelID != modelID { + descriptor.DisplayName = modelID + descriptor.Description = configuredCodexCustomDescription + } + models = append(models, descriptor) + } + return json.Marshal(struct { + Models []configuredCodexModelDescriptor `json:"models"` + }{Models: models}) +} + +func codexCatalogMetadataModels( + platform string, + modelIDs []string, + accounts []Account, + compositeRoutes []CompositeModelRoute, + compositeRoutesAvailable bool, +) map[string]string { + metadataModels := make(map[string]string, len(modelIDs)) + for _, modelID := range modelIDs { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + continue + } + metadataModelID := resolveCodexCatalogMetadataModel( + platform, + modelID, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ) + if metadataModelID != "" && metadataModelID != modelID { + metadataModels[modelID] = metadataModelID + } + } + return metadataModels +} + +func resolveCodexCatalogMetadataModel( + platform string, + modelID string, + accounts []Account, + compositeRoutes []CompositeModelRoute, + compositeRoutesAvailable bool, +) string { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + return "" + } + if platform == PlatformComposite { + if !compositeRoutesAvailable { + return modelID + } + if route, matched := matchCompositeRoute(compositeRoutes, modelID, CompositeRouteEndpointResponses); matched { + if upstreamModel := strings.TrimSpace(route.UpstreamModel); upstreamModel != "" { + return upstreamModel + } + return modelID + } + if codexCompositeRouteMatchesModel(compositeRoutes, modelID) { + return modelID + } + + claimedPlatforms := make(map[string]struct{}) + for _, account := range accounts { + accountPlatform := strings.TrimSpace(account.Platform) + if !isConcreteRequestPlatform(accountPlatform) || !codexExplicitModelMappingClaims(account, modelID) { + continue + } + claimedPlatforms[accountPlatform] = struct{}{} + } + if len(claimedPlatforms) > 1 { + return modelID + } + for accountPlatform := range claimedPlatforms { + return uniqueCodexMappedModel(accounts, accountPlatform, modelID) + } + + detectedPlatform, detected := DetectModelPlatform(modelID) + if !detected { + return modelID + } + platform = detectedPlatform + } + return uniqueCodexMappedModel(accounts, platform, modelID) +} + +func uniqueCodexMappedModel(accounts []Account, platform string, modelID string) string { + targets := make(map[string]struct{}) + for i := range accounts { + account := &accounts[i] + if account.Platform != platform { + continue + } + mappedModel, matched := account.ResolveMappedModel(modelID) + mappedModel = strings.TrimSpace(mappedModel) + if !matched || mappedModel == "" { + continue + } + targets[mappedModel] = struct{}{} + } + if len(targets) != 1 { + return modelID + } + for target := range targets { + return target + } + return modelID +} + +func groupCodexModelSupportsImageInput( + platform string, + modelID string, + accounts []Account, + compositeRoutes []CompositeModelRoute, + compositeRoutesAvailable bool, +) bool { + modelID = strings.TrimSpace(modelID) + if modelID == "" { + return false + } + upstreamModel := modelID + if platform == PlatformComposite { + var resolved bool + platform, upstreamModel, resolved = resolveCodexCompositeModelTarget( + modelID, + accounts, + compositeRoutes, + compositeRoutesAvailable, + ) + if !resolved { + return false + } + } + if platform != PlatformOpenAI && platform != PlatformGrok { + return false + } + + candidates := 0 + for i := range accounts { + account := &accounts[i] + if account.Platform != platform || !account.IsModelSupported(upstreamModel) { + continue + } + candidates++ + if !accountCodexModelSupportsImageInput(account, account.GetMappedModel(upstreamModel)) { + return false + } + } + return candidates > 0 +} + +func resolveCodexCompositeModelTarget( + modelID string, + accounts []Account, + routes []CompositeModelRoute, + routesAvailable bool, +) (string, string, bool) { + if !routesAvailable { + return "", "", false + } + if route, matched := matchCompositeRoute(routes, modelID, CompositeRouteEndpointResponses); matched { + upstreamModel := strings.TrimSpace(route.UpstreamModel) + if upstreamModel == "" { + upstreamModel = modelID + } + return route.TargetPlatform, upstreamModel, true + } + if codexCompositeRouteMatchesModel(routes, modelID) { + return "", "", false + } + + claimedPlatforms := make(map[string]struct{}) + for _, account := range accounts { + platform := strings.TrimSpace(account.Platform) + if !isConcreteRequestPlatform(platform) || !codexExplicitModelMappingClaims(account, modelID) { + continue + } + claimedPlatforms[platform] = struct{}{} + } + if len(claimedPlatforms) > 1 { + return "", "", false + } + for platform := range claimedPlatforms { + return platform, modelID, true + } + + platform, detected := DetectModelPlatform(modelID) + if !detected { + return "", "", false + } + return platform, modelID, true +} + +func codexCompositeRouteMatchesModel(routes []CompositeModelRoute, modelID string) bool { + for _, route := range routes { + publicModel := strings.TrimSpace(route.PublicModel) + if publicModel == "" { + continue + } + switch normalizeCompositeRouteMatchType(route.MatchType) { + case CompositeRouteMatchPrefix: + if strings.HasPrefix(modelID, publicModel) { + return true + } + default: + if modelID == publicModel { + return true + } + } + } + return false +} + +func codexExplicitModelMappingClaims(account Account, modelID string) bool { + if account.Credentials == nil || strings.TrimSpace(modelID) == "" { + return false + } + mapped := strings.TrimSpace(account.GetModelMapping()[modelID]) + return mapped != "" +} + +func accountCodexModelSupportsImageInput(account *Account, upstreamModel string) bool { + if account == nil { + return false + } + switch account.Platform { + case PlatformOpenAI: + if !isOpenAICodexImageInputModel(upstreamModel) { + return false + } + if account.IsOpenAIOAuth() { + return true + } + if !account.IsOpenAIApiKey() { + return false + } + baseURL := strings.TrimSpace(account.GetCredential("base_url")) + return baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) + case PlatformGrok: + if !isOfficialGrokCodexBaseURL(account.GetGrokBaseURL()) { + return false + } + canonical := xai.ResolveGrokTextResponsesModelID(upstreamModel) + return isGrokCodexImageInputModel(canonical) + default: + return false + } +} + +func isGrokCodexImageInputModel(model string) bool { + switch strings.ToLower(strings.TrimSpace(model)) { + case "grok-4.3", + "grok-4.5", + "grok-4.6", + "grok-build-0.1", + "grok-4.20-0309-reasoning", + "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent-0309": + return true + default: + return false + } +} + +func isOfficialGrokCodexBaseURL(raw string) bool { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || parsed.Host == "" { + return false + } + return xai.IsOfficialBaseURLHost(strings.TrimSuffix(parsed.Hostname(), ".")) +} + +// BuildDeepSeekCodexModelsManifest preserves the historical entry point for +// callers that still use the provider-specific function name. +func BuildDeepSeekCodexModelsManifest(modelIDs []string) ([]byte, error) { + return BuildCodexModelsManifest(modelIDs) +} + +func mergeConfiguredCodexModelsManifest( + body []byte, + configuredModels []string, + selectedModels []string, + filterBySelection bool, +) ([]byte, bool, error) { + var envelope map[string]json.RawMessage + if err := json.Unmarshal(body, &envelope); err != nil { + return nil, false, err + } + var upstreamModels []json.RawMessage + if err := json.Unmarshal(envelope["models"], &upstreamModels); err != nil { + return nil, false, err + } + + selected := make(map[string]struct{}, len(selectedModels)) + for _, modelID := range selectedModels { + modelID = strings.TrimSpace(modelID) + if modelID != "" { + selected[modelID] = struct{}{} + } + } + seen := make(map[string]struct{}, len(upstreamModels)+len(configuredModels)) + merged := make([]json.RawMessage, 0, len(upstreamModels)+len(configuredModels)) + changed := false + for _, rawModel := range upstreamModels { + var descriptor struct { + Slug string `json:"slug"` + } + if err := json.Unmarshal(rawModel, &descriptor); err != nil || strings.TrimSpace(descriptor.Slug) == "" { + if filterBySelection { + changed = true + continue + } + merged = append(merged, rawModel) + continue + } + descriptor.Slug = strings.TrimSpace(descriptor.Slug) + if isCodexDedicatedMediaModel(descriptor.Slug) { + changed = true + continue + } + if filterBySelection { + if _, allowed := selected[descriptor.Slug]; !allowed { + changed = true + continue + } + } + if strings.HasPrefix(descriptor.Slug, codexAutoModelPrefix) { + _, explicitlyEnabled := selected[descriptor.Slug] + explicitlyEnabled = filterBySelection && explicitlyEnabled + if !explicitlyEnabled { + changed = true + continue + } + visibleModel, visibilityChanged, err := codexModelWithVisibility(rawModel, "list") + if err != nil { + return nil, false, err + } + rawModel = visibleModel + changed = changed || visibilityChanged + } + seen[descriptor.Slug] = struct{}{} + merged = append(merged, rawModel) + } + + for _, modelID := range configuredModels { + if isCodexDedicatedMediaModel(modelID) { + continue + } + if filterBySelection { + if _, allowed := selected[modelID]; !allowed { + continue + } + } + if strings.HasPrefix(modelID, codexAutoModelPrefix) { + if _, explicitlyEnabled := selected[modelID]; !filterBySelection || !explicitlyEnabled { + continue + } + } + if _, exists := seen[modelID]; exists { + continue + } + rawModel, err := json.Marshal(newConfiguredCodexModelDescriptor(modelID)) + if err != nil { + return nil, false, err + } + merged = append(merged, rawModel) + seen[modelID] = struct{}{} + changed = true + } + if !changed { + return body, false, nil + } + + rawModels, err := json.Marshal(merged) + if err != nil { + return nil, false, err + } + envelope["models"] = rawModels + mergedBody, err := json.Marshal(envelope) + if err != nil { + return nil, false, err + } + return mergedBody, true, nil +} + +func codexModelWithVisibility(rawModel json.RawMessage, visibility string) (json.RawMessage, bool, error) { + var fields map[string]json.RawMessage + if err := json.Unmarshal(rawModel, &fields); err != nil { + return nil, false, err + } + var current string + if rawVisibility, ok := fields["visibility"]; ok { + if err := json.Unmarshal(rawVisibility, ¤t); err == nil && current == visibility { + return rawModel, false, nil + } + } + rawVisibility, err := json.Marshal(visibility) + if err != nil { + return nil, false, err + } + fields["visibility"] = rawVisibility + updated, err := json.Marshal(fields) + if err != nil { + return nil, false, err + } + return updated, true, nil } type codexModelsManifestUpstreamError struct { @@ -55,12 +1235,13 @@ 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 and ChatGPT-backend 401, are intentionally not -// retried. A manifest 401 from the ChatGPT Codex backend reflects the selected -// OAuth account's upstream token rather than the client request (the client's -// own API key was already validated locally), so a different account may still -// serve the manifest. Custom API key upstreams keep the old no-failover 401 +// may succeed without changing the request. API key upstream 404/405 responses +// mean that the selected account does not expose a model-discovery endpoint, so +// another account may still serve the manifest. Other upstream 4xx responses, +// except 429 and ChatGPT-backend 401, are intentionally not retried. A manifest +// 401 from the ChatGPT Codex backend reflects the selected OAuth account's +// upstream token rather than the client request (the client's own API key was +// already validated locally). Custom API key upstreams keep the no-failover 401 // behavior because their /models auth semantics are not authoritative for the // account. func IsRetryableCodexModelsManifestError(err error) bool { @@ -68,6 +1249,14 @@ func IsRetryableCodexModelsManifestError(err error) bool { return errors.As(err, &upstreamErr) && upstreamErr.retryable } +func isRetryableCodexModelsManifestStatus(statusCode int, useAPIKeyUpstream bool) bool { + return (useAPIKeyUpstream && + (statusCode == http.StatusNotFound || statusCode == http.StatusMethodNotAllowed)) || + (statusCode == http.StatusUnauthorized && !useAPIKeyUpstream) || + statusCode == http.StatusTooManyRequests || + (statusCode >= http.StatusInternalServerError && statusCode < 600) +} + func isRetryableCodexModelsManifestTransportError(err error) bool { if err == nil || errors.Is(err, context.Canceled) { return false @@ -193,6 +1382,10 @@ func (c *codexModelsManifestCache) set(key string, manifest *CodexModelsManifest if manifest == nil || len(manifest.Body) > codexModelsManifestCacheBodyLimit { return } + remainingBodyBudget := codexModelsManifestCacheBodyLimit - len(manifest.Body) + if len(manifest.upstreamSourceBody) > remainingBodyBudget { + return + } c.mu.Lock() defer c.mu.Unlock() if c.entries == nil { @@ -255,14 +1448,7 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc 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", - ) - } + baseURL := strings.TrimSpace(credAccount.GetOpenAIBaseURL()) 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") @@ -508,9 +1694,7 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont statusCode: resp.StatusCode, headers: resp.Header.Clone(), body: body, - retryable: (resp.StatusCode == http.StatusUnauthorized && !request.useAPIKeyUpstream) || - resp.StatusCode == http.StatusTooManyRequests || - (resp.StatusCode >= http.StatusInternalServerError && resp.StatusCode < 600), + retryable: isRetryableCodexModelsManifestStatus(resp.StatusCode, request.useAPIKeyUpstream), } } @@ -529,8 +1713,11 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont } } upstreamBody := body + convertedFromOpenAIModelList := false if request.useAPIKeyUpstream { - body = convertOpenAIModelListToCodexManifest(body) + convertedBody := convertOpenAIModelListToCodexManifest(body) + convertedFromOpenAIModelList = !bytes.Equal(convertedBody, body) + body = convertedBody } if err := validateCodexModelsManifestEnvelope(body); err != nil { return nil, &codexModelsManifestUpstreamError{ @@ -544,6 +1731,22 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont } } if request.useAPIKeyUpstream { + body, err = completeAPIKeyCodexModelsManifestMetadata( + body, + false, + request.credentialAccount != nil && isOfficialOpenAIModelsBaseURL(request.credentialAccount.GetOpenAIBaseURL()), + ) + if err != nil { + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf( + http.StatusBadGateway, + "OPENAI_CODEX_MODELS_UPSTREAM_INVALID_MANIFEST", + "codex models manifest upstream metadata could not be completed: %v", + err, + ), + retryable: true, + } + } body, err = adjustAPIKeyCodexModelsManifest(body) if err != nil { return nil, &codexModelsManifestUpstreamError{ @@ -558,7 +1761,12 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont } } etag := resp.Header.Get("ETag") - manifest := &CodexModelsManifest{Body: body, ETag: etag} + manifest := &CodexModelsManifest{ + Body: body, + ETag: etag, + upstreamSourceBody: append([]byte(nil), upstreamBody...), + convertedFromOpenAIModelList: convertedFromOpenAIModelList, + } if request.useAPIKeyUpstream { manifest.upstreamETag = etag if !bytes.Equal(body, upstreamBody) { @@ -573,6 +1781,12 @@ func codexModelsManifestBodyETag(body []byte) string { return fmt.Sprintf(`"%x"`, sum) } +// CodexModelsManifestETag returns the strong ETag for a generated client +// catalog. It is based on the final JSON body after local filtering. +func CodexModelsManifestETag(body []byte) string { + return codexModelsManifestBodyETag(body) +} + var apiKeyCodexModelsWithoutResponsesLite = map[string]struct{}{ "gpt-5.6-sol": {}, "gpt-5.6-terra": {}, @@ -636,9 +1850,8 @@ func adjustAPIKeyCodexModelsManifest(body []byte) ([]byte, error) { // convertOpenAIModelListToCodexManifest rewrites a standard OpenAI // GET /v1/models response ({"object":"list","data":[{"id":...},...]}) into the -// Codex manifest envelope ({"models":[{"slug":...},...]}) so custom API key -// upstreams that only implement the standard endpoint can serve Codex model -// discovery. Bodies that already carry a top-level models field, are not the +// same complete Codex manifest used by locally generated custom-provider +// catalogs. Bodies that already carry a top-level models field, are not the // standard list shape, or yield no usable model IDs are returned unchanged so // envelope validation reports the original payload. func convertOpenAIModelListToCodexManifest(body []byte) []byte { @@ -659,27 +1872,316 @@ func convertOpenAIModelListToCodexManifest(body []byte) []byte { if err := json.Unmarshal(data, &entries); err != nil { return body } - type codexModelEntry struct { - Slug string `json:"slug"` - } - models := make([]codexModelEntry, 0, len(entries)) + modelIDs := make([]string, 0, len(entries)) for _, entry := range entries { id := strings.TrimSpace(entry.ID) if id == "" { continue } - models = append(models, codexModelEntry{Slug: id}) + modelIDs = append(modelIDs, id) } - if len(models) == 0 { + if len(modelIDs) == 0 { return body } - converted, err := json.Marshal(map[string][]codexModelEntry{"models": models}) + converted, err := BuildCodexModelsManifest(modelIDs) if err != nil { return body } return converted } +// completeAPIKeyCodexModelsManifestMetadata fills fields omitted by standard +// OpenAI-compatible /models endpoints. Existing provider metadata always wins; +// only absent or null values are synthesized. +// CompleteAPIKeyCodexModelsManifestForClient fills the complete ModelInfo +// contract immediately before a group-specific API key manifest is returned. +// The shared upstream cache remains independent from local group policy. +func (s *OpenAIGatewayService) CompleteAPIKeyCodexModelsManifestForClient(manifest *CodexModelsManifest, account *Account) error { + if manifest == nil || account == nil || !account.IsOpenAIApiKey() || manifest.NotModified || len(manifest.Body) == 0 { + return nil + } + body := manifest.Body + if len(manifest.upstreamSourceBody) > 0 { + body = append([]byte(nil), manifest.upstreamSourceBody...) + if manifest.convertedFromOpenAIModelList { + body = convertOpenAIModelListToCodexManifest(body) + } + } + var err error + body, err = applySyncedAPIKeyCodexModelMetadata(body, account, manifest.convertedFromOpenAIModelList) + if err != nil { + return err + } + body, err = completeAPIKeyCodexModelsManifestMetadata( + body, + true, + isOfficialOpenAIModelsBaseURL(account.GetOpenAIBaseURL()), + ) + if err != nil { + return err + } + body, err = adjustAPIKeyCodexModelsManifest(body) + if err != nil { + return err + } + manifest.Body = body + manifest.ETag = codexModelsManifestBodyETag(manifest.Body) + return nil +} + +func applySyncedAPIKeyCodexModelMetadata(body []byte, account *Account, overwriteLocalDefaults bool) ([]byte, error) { + snapshot := account.GetUpstreamModelMetadataSnapshot() + if snapshot == nil || len(snapshot.Models) == 0 { + return body, nil + } + + var envelope map[string]json.RawMessage + if err := json.Unmarshal(body, &envelope); err != nil { + return nil, fmt.Errorf("decode JSON object: %w", err) + } + var models []json.RawMessage + if err := json.Unmarshal(envelope["models"], &models); err != nil { + return nil, fmt.Errorf("decode top-level models array: %w", err) + } + + changed := false + for i, rawModel := range models { + var model map[string]json.RawMessage + if err := json.Unmarshal(rawModel, &model); err != nil || model == nil { + continue + } + var slug string + if err := json.Unmarshal(model["slug"], &slug); err != nil { + continue + } + slug = strings.TrimSpace(slug) + metadata, ok := snapshot.Models[slug] + if !ok { + continue + } + + descriptor := newConfiguredCodexModelDescriptor(slug) + applyUpstreamModelMetadataToCodexDescriptor( + &descriptor, + codexModelMetadataOverride{UpstreamModelMetadata: metadata}, + ) + descriptorBody, err := json.Marshal(descriptor) + if err != nil { + return nil, fmt.Errorf("encode synced model %q: %w", slug, err) + } + var syncedFields map[string]json.RawMessage + if err := json.Unmarshal(descriptorBody, &syncedFields); err != nil { + return nil, fmt.Errorf("decode synced model %q: %w", slug, err) + } + + fields := make([]string, 0, 7) + if strings.TrimSpace(metadata.DisplayName) != "" { + fields = append(fields, "display_name") + } + if strings.TrimSpace(metadata.Description) != "" { + fields = append(fields, "description") + } + if metadata.Reasoning != nil { + fields = append(fields, "default_reasoning_level", "supported_reasoning_levels") + } + if len(normalizeCodexInputModalities(metadata.InputModalities)) > 0 { + fields = append(fields, "input_modalities") + } + if metadata.ContextWindow > 0 { + fields = append(fields, "context_window", "max_context_window") + } + + modelChanged := false + for _, field := range fields { + value, exists := syncedFields[field] + if !exists { + continue + } + current, currentExists := model[field] + current = bytes.TrimSpace(current) + if !overwriteLocalDefaults && currentExists && len(current) > 0 && !bytes.Equal(current, []byte("null")) { + continue + } + if bytes.Equal(current, bytes.TrimSpace(value)) { + continue + } + model[field] = value + modelChanged = true + } + if !modelChanged { + continue + } + encoded, err := json.Marshal(model) + if err != nil { + return nil, fmt.Errorf("encode model %q with synced metadata: %w", slug, err) + } + models[i] = encoded + changed = true + } + if !changed { + return body, nil + } + + encodedModels, err := json.Marshal(models) + if err != nil { + return nil, fmt.Errorf("encode models with synced metadata: %w", err) + } + envelope["models"] = encodedModels + updated, err := json.Marshal(envelope) + if err != nil { + return nil, fmt.Errorf("encode manifest with synced metadata: %w", err) + } + return updated, nil +} + +func completeAPIKeyCodexModelsManifestMetadata(body []byte, completeAll, officialOpenAI bool) ([]byte, error) { + var envelope map[string]json.RawMessage + if err := json.Unmarshal(body, &envelope); err != nil { + return nil, fmt.Errorf("decode JSON object: %w", err) + } + var models []json.RawMessage + if err := json.Unmarshal(envelope["models"], &models); err != nil { + return nil, fmt.Errorf("decode top-level models array: %w", err) + } + + changed := false + if officialOpenAI { + filtered := make([]json.RawMessage, 0, len(models)) + for _, rawModel := range models { + var model struct { + Slug string `json:"slug"` + } + if err := json.Unmarshal(rawModel, &model); err != nil || strings.TrimSpace(model.Slug) == "" { + filtered = append(filtered, rawModel) + continue + } + if !isOfficialOpenAICodexCatalogModel(model.Slug) { + changed = true + continue + } + filtered = append(filtered, rawModel) + } + models = filtered + } + for i, rawModel := range models { + var model map[string]json.RawMessage + if err := json.Unmarshal(rawModel, &model); err != nil || model == nil { + continue + } + var slug string + if err := json.Unmarshal(model["slug"], &slug); err != nil { + continue + } + slug = strings.TrimSpace(slug) + if slug == "" { + continue + } + + completeDescriptor := completeAll || isDeepSeekCodexModel(slug) + forceOfficialImage := officialOpenAI && isOpenAICodexImageInputModel(slug) + if !completeDescriptor && !forceOfficialImage { + continue + } + + descriptor := newConfiguredCodexModelDescriptor(slug) + if forceOfficialImage { + descriptor.InputModalities = []string{"text", "image"} + descriptor.SupportsImageDetailOriginal = true + } + defaultBody, err := json.Marshal(descriptor) + if err != nil { + return nil, fmt.Errorf("encode default model %q: %w", slug, err) + } + var defaults map[string]json.RawMessage + if err := json.Unmarshal(defaultBody, &defaults); err != nil { + return nil, fmt.Errorf("decode default model %q: %w", slug, err) + } + + modelChanged := false + if completeDescriptor { + merged, err := mergeMissingCodexModelFields(model, defaults) + if err != nil { + return nil, fmt.Errorf("complete model %q: %w", slug, err) + } + modelChanged = merged + } + if forceOfficialImage { + modalities, err := json.Marshal([]string{"text", "image"}) + if err != nil { + return nil, fmt.Errorf("encode input modalities for model %q: %w", slug, err) + } + if !bytes.Equal(bytes.TrimSpace(model["input_modalities"]), modalities) { + model["input_modalities"] = modalities + modelChanged = true + } + imageDetailOriginal := json.RawMessage("true") + if !bytes.Equal(bytes.TrimSpace(model["supports_image_detail_original"]), imageDetailOriginal) { + model["supports_image_detail_original"] = imageDetailOriginal + modelChanged = true + } + } + if !modelChanged { + continue + } + encoded, err := json.Marshal(model) + if err != nil { + return nil, fmt.Errorf("encode completed model %q: %w", slug, err) + } + models[i] = encoded + changed = true + } + if !changed { + return body, nil + } + + encodedModels, err := json.Marshal(models) + if err != nil { + return nil, fmt.Errorf("encode top-level models array: %w", err) + } + envelope["models"] = encodedModels + completed, err := json.Marshal(envelope) + if err != nil { + return nil, fmt.Errorf("encode JSON object: %w", err) + } + return completed, nil +} + +func mergeMissingCodexModelFields(current, defaults map[string]json.RawMessage) (bool, error) { + changed := false + for key, defaultValue := range defaults { + currentValue, exists := current[key] + if !exists || (bytes.Equal(bytes.TrimSpace(currentValue), []byte("null")) && + !bytes.Equal(bytes.TrimSpace(defaultValue), []byte("null"))) { + current[key] = defaultValue + changed = true + continue + } + + var currentObject map[string]json.RawMessage + var defaultObject map[string]json.RawMessage + if err := json.Unmarshal(currentValue, ¤tObject); err != nil || currentObject == nil { + continue + } + if err := json.Unmarshal(defaultValue, &defaultObject); err != nil || defaultObject == nil { + continue + } + nestedChanged, err := mergeMissingCodexModelFields(currentObject, defaultObject) + if err != nil { + return false, err + } + if !nestedChanged { + continue + } + mergedValue, err := json.Marshal(currentObject) + if err != nil { + return false, fmt.Errorf("encode field %q: %w", key, err) + } + current[key] = mergedValue + changed = true + } + return changed, nil +} + func validateCodexModelsManifestEnvelope(body []byte) error { var envelope map[string]json.RawMessage if err := json.Unmarshal(body, &envelope); err != nil { @@ -720,6 +2222,20 @@ func buildCodexModelsManifestCacheKey(request codexModelsManifestRequest) string return fmt.Sprintf("%x", hasher.Sum(nil)) } +func cloneCodexModelsManifest(manifest *CodexModelsManifest) *CodexModelsManifest { + if manifest == nil { + return nil + } + cloned := *manifest + if manifest.Body != nil { + cloned.Body = append([]byte(nil), manifest.Body...) + } + if manifest.upstreamSourceBody != nil { + cloned.upstreamSourceBody = append([]byte(nil), manifest.upstreamSourceBody...) + } + return &cloned +} + func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch string) *CodexModelsManifest { if manifest == nil { return nil @@ -727,7 +2243,7 @@ func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch str if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { return &CodexModelsManifest{ETag: manifest.ETag, NotModified: true} } - return manifest + return cloneCodexModelsManifest(manifest) } func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool { @@ -752,6 +2268,12 @@ func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool { return false } +// CodexModelsManifestETagMatches applies If-None-Match semantics to a Codex +// catalog ETag, including weak and comma-separated validators. +func CodexModelsManifestETagMatches(ifNoneMatch, etag string) bool { + return codexModelsManifestETagMatches(ifNoneMatch, etag) +} + func isOfficialOpenAIModelsBaseURL(raw string) bool { parsed, err := url.Parse(strings.TrimSpace(raw)) if err != nil { diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index 503d494857..77a65cfb38 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -3,6 +3,7 @@ package service import ( "bytes" "context" + "encoding/json" "errors" "fmt" "io" @@ -28,6 +29,1165 @@ type codexModelsHTTPUpstreamStub struct { do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) } +type codexModelsVisibilityAccountRepo struct { + AccountRepository + byGroup map[int64][]Account +} + +func (r codexModelsVisibilityAccountRepo) ListSchedulableByGroupID(_ context.Context, groupID int64) ([]Account, error) { + accounts := r.byGroup[groupID] + return append([]Account(nil), accounts...), nil +} + +func (r codexModelsVisibilityAccountRepo) ListByGroup(_ context.Context, groupID int64) ([]Account, error) { + accounts := r.byGroup[groupID] + return append([]Account(nil), accounts...), nil +} + +type countingCodexModelsAccountRepo struct { + AccountRepository + accounts []Account + err error + listByGroupErr error + calls atomic.Int32 +} + +func (r *countingCodexModelsAccountRepo) ListSchedulableByGroupID(_ context.Context, _ int64) ([]Account, error) { + r.calls.Add(1) + if r.err != nil { + return nil, r.err + } + return append([]Account(nil), r.accounts...), nil +} + +func (r *countingCodexModelsAccountRepo) ListByGroup(_ context.Context, _ int64) ([]Account, error) { + if r.listByGroupErr != nil { + return nil, r.listByGroupErr + } + return append([]Account(nil), r.accounts...), nil +} + +type splitCodexModelsAccountRepo struct { + AccountRepository + schedulable map[int64][]Account + catalog map[int64][]Account +} + +func (r splitCodexModelsAccountRepo) ListSchedulableByGroupID(_ context.Context, groupID int64) ([]Account, error) { + return append([]Account(nil), r.schedulable[groupID]...), nil +} + +func (r splitCodexModelsAccountRepo) ListByGroup(_ context.Context, groupID int64) ([]Account, error) { + return append([]Account(nil), r.catalog[groupID]...), nil +} + +func newCodexCatalogMappedAccount( + id int64, + target string, + displayName string, + levels []string, + modalities []string, + contextWindow int64, + schedulable bool, + extraMapping map[string]any, +) Account { + reasoning := true + mapping := map[string]any{"my-coder": target} + for key, value := range extraMapping { + mapping[key] = value + } + account := Account{ + ID: id, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: schedulable, + Credentials: map[string]any{ + "base_url": fmt.Sprintf("https://provider-%d.example/v1", id), + "model_mapping": mapping, + }, + } + models := map[string]UpstreamModelMetadata{ + target: { + ID: target, + DisplayName: displayName, + Description: displayName + " upstream", + Reasoning: &reasoning, + SupportedReasoningLevels: levels, + InputModalities: modalities, + ContextWindow: contextWindow, + }, + } + for _, value := range extraMapping { + exclusive, _ := value.(string) + if exclusive == "" || exclusive == target { + continue + } + models[exclusive] = UpstreamModelMetadata{ + ID: exclusive, + DisplayName: "Exclusive Model", + Description: "Only mapped on the unschedulable account", + Reasoning: &reasoning, + SupportedReasoningLevels: []string{"high"}, + InputModalities: []string{"text", "image"}, + ContextWindow: 1_000_000, + } + } + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: models}) + return account +} + +func TestFilterCodexModelIDsForGroupOmitsWildcardKeys(t *testing.T) { + t.Parallel() + + got := FilterCodexModelIDsForGroup( + []string{"deepseek-v4-pro", "foo-*", " bar-* ", "gpt-5.5"}, + &Group{Platform: PlatformDeepseek}, + ) + require.Equal(t, []string{"deepseek-v4-pro", "gpt-5.5"}, got) +} + +func decodeCodexManifestModels(t *testing.T, body []byte) []map[string]any { + t.Helper() + + var envelope struct { + Models []map[string]any `json:"models"` + } + require.NoError(t, json.Unmarshal(body, &envelope)) + return envelope.Models +} + +func codexManifestModelSlugs(t *testing.T, body []byte) []string { + t.Helper() + + models := decodeCodexManifestModels(t, body) + slugs := make([]string, 0, len(models)) + for _, model := range models { + slug, ok := model["slug"].(string) + require.True(t, ok) + slugs = append(slugs, slug) + } + return slugs +} + +func requireCompleteConfiguredCodexModel(t *testing.T, model map[string]any, slug string) { + t.Helper() + + require.Equal(t, slug, model["slug"]) + require.NotEmpty(t, model["display_name"]) + require.NotEmpty(t, model["description"]) + require.Equal(t, "unified_exec", model["shell_type"]) + require.Equal(t, "list", model["visibility"]) + require.Equal(t, true, model["supported_in_api"]) + require.NotNil(t, model["priority"]) + require.Equal(t, []any{}, model["additional_speed_tiers"]) + require.Equal(t, []any{}, model["service_tiers"]) + require.Contains(t, model, "default_service_tier") + require.Contains(t, model, "availability_nux") + require.Contains(t, model, "upgrade") + require.Contains(t, model, "default_verbosity") + require.Contains(t, model, "apply_patch_tool_type") + require.Contains(t, model, "auto_compact_token_limit") + require.Contains(t, model, "comp_hash") + require.Contains(t, model, "auto_review_model_override") + require.Contains(t, model, "model_specialty") + require.Contains(t, model, "tool_mode") + require.Contains(t, model, "multi_agent_version") + require.Equal(t, true, model["supports_reasoning_summary_parameter"]) + require.Contains(t, model, "include_skills_usage_instructions") + require.Contains(t, model, "include_plugin_usage_instructions") + require.Contains(t, model, "include_apps_usage_instructions") + require.Contains(t, model, "supports_image_detail_original") + require.Contains(t, model, "node_repl_auto_review_required") + require.Contains(t, model, "node_repl_disabled") + require.Contains(t, model, "truncation_policy") + require.Contains(t, model, "supports_parallel_tool_calls") + require.Contains(t, model, "experimental_supported_tools") + modelMessages, ok := model["model_messages"].(map[string]any) + require.True(t, ok) + require.NotEmpty(t, modelMessages["instructions_template"]) + for _, key := range []string{ + "instructions_variables", + "approvals", + "collaboration_modes", + "auto_review", + "permissions", + "multi_agent", + "token_budget", + "guardian_v2", + } { + require.Contains(t, modelMessages, key) + } +} + +func effortsFromManifestModel(t *testing.T, model map[string]any) []string { + t.Helper() + + levels, ok := model["supported_reasoning_levels"].([]any) + require.True(t, ok) + efforts := make([]string, 0, len(levels)) + for _, rawLevel := range levels { + level, ok := rawLevel.(map[string]any) + require.True(t, ok) + effort, ok := level["effort"].(string) + require.True(t, ok) + efforts = append(efforts, effort) + } + return efforts +} + +func TestNewConfiguredCodexModelDescriptorUsesProviderMetadataAndSafeFallback(t *testing.T) { + t.Parallel() + + deepSeek := newConfiguredCodexModelDescriptor("deepseek-v4-pro") + require.Equal(t, "DeepSeek V4 Pro", deepSeek.DisplayName) + require.Equal(t, int64(1_000_000), deepSeek.ContextWindow) + require.Equal(t, int64(1_000_000), deepSeek.MaxContextWindow) + require.NotNil(t, deepSeek.DefaultReasoningLevel) + require.Equal(t, "high", *deepSeek.DefaultReasoningLevel) + require.Equal(t, []configuredCodexReasoningLevel{ + {Effort: "low", Description: "Fast responses with lighter reasoning"}, + {Effort: "high", Description: "Greater reasoning depth for coding and agent tasks"}, + {Effort: "max", Description: "Maximum reasoning depth for complex tasks"}, + }, deepSeek.SupportedReasoningLevels) + require.True(t, deepSeek.SupportsParallelToolCalls) + require.Equal(t, []string{"text"}, deepSeek.InputModalities) + + grok := newConfiguredCodexModelDescriptor("grok-4.6") + require.Equal(t, "Grok 4.6", grok.DisplayName) + require.Equal(t, int64(500_000), grok.ContextWindow) + require.Equal(t, int64(500_000), grok.MaxContextWindow) + require.NotNil(t, grok.DefaultReasoningLevel) + require.Equal(t, "high", *grok.DefaultReasoningLevel) + require.Equal(t, []string{"low", "medium", "high", "xhigh"}, effortsFromConfiguredCodexLevels(grok.SupportedReasoningLevels)) + require.True(t, grok.SupportsParallelToolCalls) + require.Equal(t, []string{"text"}, grok.InputModalities) + require.NotContains(t, grok.SupportedReasoningLevels, configuredCodexReasoningLevel{Effort: "none"}) + require.NotContains(t, grok.SupportedReasoningLevels, configuredCodexReasoningLevel{Effort: "max"}) + + grokAlias := newConfiguredCodexModelDescriptor("xai/grok-4.6-latest") + require.Equal(t, "Grok 4.6", grokAlias.DisplayName) + require.NotNil(t, grokAlias.DefaultReasoningLevel) + require.Equal(t, "high", *grokAlias.DefaultReasoningLevel) + require.Equal(t, []string{"low", "medium", "high", "xhigh"}, effortsFromConfiguredCodexLevels(grokAlias.SupportedReasoningLevels)) + + grok45 := newConfiguredCodexModelDescriptor("grok-4.5") + require.Equal(t, []string{"low", "medium", "high"}, effortsFromConfiguredCodexLevels(grok45.SupportedReasoningLevels)) + + grokNonReasoning := newConfiguredCodexModelDescriptor("grok-4.20-0309-non-reasoning") + require.Equal(t, "Grok 4.20 Non Reasoning", grokNonReasoning.DisplayName) + require.NotNil(t, grokNonReasoning.DefaultReasoningLevel) + require.Equal(t, "none", *grokNonReasoning.DefaultReasoningLevel) + require.Equal(t, []configuredCodexReasoningLevel{ + {Effort: "none", Description: configuredCodexReasoningLevelDescription("none")}, + }, grokNonReasoning.SupportedReasoningLevels) + + claude := newConfiguredCodexModelDescriptor("claude-opus-4-6") + require.Equal(t, "Claude Opus 4.6", claude.DisplayName) + require.NotNil(t, claude.DefaultReasoningLevel) + require.Equal(t, "medium", *claude.DefaultReasoningLevel) + require.Equal(t, []string{"low", "medium", "high", "max"}, effortsFromConfiguredCodexLevels(claude.SupportedReasoningLevels)) + require.NotContains(t, claude.SupportedReasoningLevels, configuredCodexReasoningLevel{Effort: "xhigh"}) + require.NotContains(t, claude.SupportedReasoningLevels, configuredCodexReasoningLevel{Effort: "none"}) + require.True(t, claude.SupportsParallelToolCalls) + + claudeOpus5 := newConfiguredCodexModelDescriptor("claude-opus-5") + require.Equal(t, []string{"low", "medium", "high", "xhigh", "max"}, effortsFromConfiguredCodexLevels(claudeOpus5.SupportedReasoningLevels)) + + providerQualifiedClaude := newConfiguredCodexModelDescriptor("anthropic/claude-sonnet-4-6") + require.Equal(t, "Claude Sonnet 4.6", providerQualifiedClaude.DisplayName) + require.Equal(t, []string{"low", "medium", "high", "max"}, effortsFromConfiguredCodexLevels(providerQualifiedClaude.SupportedReasoningLevels)) + + claudeHaiku := newConfiguredCodexModelDescriptor("claude-haiku-4-5-20251001") + require.Equal(t, "Claude Haiku 4.5", claudeHaiku.DisplayName) + require.NotNil(t, claudeHaiku.DefaultReasoningLevel) + require.Equal(t, "none", *claudeHaiku.DefaultReasoningLevel) + require.Equal(t, []string{"none"}, effortsFromConfiguredCodexLevels(claudeHaiku.SupportedReasoningLevels)) + + gpt56 := newConfiguredCodexModelDescriptor("gpt-5.6-sol") + require.Equal(t, "GPT-5.6 Sol", gpt56.DisplayName) + require.Equal(t, "OpenAI GPT coding model routed through Sub2API.", gpt56.Description) + require.NotNil(t, gpt56.DefaultReasoningLevel) + require.Equal(t, "low", *gpt56.DefaultReasoningLevel) + require.Equal(t, configuredCodexGPTReasoningLevels("gpt-5.6-sol"), gpt56.SupportedReasoningLevels) + require.Equal(t, []string{"low", "medium", "high", "xhigh", "max", "ultra"}, effortsFromConfiguredCodexLevels(gpt56.SupportedReasoningLevels)) + require.True(t, gpt56.SupportsParallelToolCalls) + require.True(t, gpt56.SupportVerbosity) + require.Equal(t, []string{"text"}, gpt56.InputModalities) + require.Equal(t, int64(872_000), gpt56.MaxContextWindow) + require.Equal(t, configuredCodexTruncationPolicy{Mode: "tokens", Limit: 10_000}, gpt56.TruncationPolicy) + require.NotNil(t, gpt56.DefaultVerbosity) + require.Equal(t, "low", *gpt56.DefaultVerbosity) + require.True(t, gpt56.SupportsReasoningSummaryParameter) + require.Equal(t, "none", gpt56.DefaultReasoningSummary) + + gpt56Luna := newConfiguredCodexModelDescriptor("gpt-5.6-luna") + require.Equal(t, []string{"low", "medium", "high", "xhigh", "max"}, effortsFromConfiguredCodexLevels(gpt56Luna.SupportedReasoningLevels)) + require.Equal(t, "medium", *gpt56Luna.DefaultReasoningLevel) + + gpt55 := newConfiguredCodexModelDescriptor("gpt-5.5") + require.Equal(t, "GPT-5.5", gpt55.DisplayName) + require.NotNil(t, gpt55.DefaultReasoningLevel) + require.Equal(t, "medium", *gpt55.DefaultReasoningLevel) + require.Equal(t, []string{"low", "medium", "high", "xhigh"}, effortsFromConfiguredCodexLevels(gpt55.SupportedReasoningLevels)) + require.NotContains(t, gpt55.SupportedReasoningLevels, configuredCodexReasoningLevel{Effort: "max"}) + require.NotNil(t, gpt55.DefaultVerbosity) + require.Equal(t, "low", *gpt55.DefaultVerbosity) + + gpt4o := newConfiguredCodexModelDescriptor("gpt-4o") + require.Equal(t, "gpt-4o", gpt4o.DisplayName) + require.NotNil(t, gpt4o.DefaultReasoningLevel) + require.Equal(t, "none", *gpt4o.DefaultReasoningLevel) + require.Equal(t, []string{"none"}, effortsFromConfiguredCodexLevels(gpt4o.SupportedReasoningLevels)) + require.True(t, gpt4o.SupportsParallelToolCalls) + + image := newConfiguredCodexModelDescriptor("gpt-image-2") + require.Equal(t, "gpt-image-2", image.DisplayName) + require.NotNil(t, image.DefaultReasoningLevel) + require.Equal(t, "none", *image.DefaultReasoningLevel) + require.Equal(t, []string{"none"}, effortsFromConfiguredCodexLevels(image.SupportedReasoningLevels)) + + custom := newConfiguredCodexModelDescriptor("company-coding-model") + require.Equal(t, "company-coding-model", custom.DisplayName) + require.Equal(t, int64(272_000), custom.ContextWindow) + require.NotNil(t, custom.DefaultReasoningLevel) + require.Equal(t, "none", *custom.DefaultReasoningLevel) + require.Equal(t, []configuredCodexReasoningLevel{ + {Effort: "none", Description: configuredCodexReasoningLevelDescription("none")}, + }, custom.SupportedReasoningLevels) + require.False(t, custom.SupportsParallelToolCalls) + require.NotEmpty(t, custom.ModelMessages.InstructionsTemplate) + require.Equal(t, "auto", custom.DefaultReasoningSummary) + require.Equal(t, configuredCodexTruncationPolicy{Mode: "bytes", Limit: 10_000}, custom.TruncationPolicy) +} + +func effortsFromConfiguredCodexLevels(levels []configuredCodexReasoningLevel) []string { + efforts := make([]string, 0, len(levels)) + for _, level := range levels { + efforts = append(efforts, level.Effort) + } + return efforts +} + +// Scenario: 无推理模型可直接选中。 +func TestBuildCodexModelsManifestUsesSingleNoneReasoningChoiceForCustomModel(t *testing.T) { + t.Parallel() + + body, err := BuildCodexModelsManifest([]string{"company-coding-model"}) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "none", models[0]["default_reasoning_level"]) + levels, ok := models[0]["supported_reasoning_levels"].([]any) + require.True(t, ok) + require.Len(t, levels, 1) + firstLevel, ok := levels[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "none", firstLevel["effort"]) +} + +// Scenario: 已知推理模型保留真实档位。 +func TestBuildCodexModelsManifestKeepsKnownReasoningChoices(t *testing.T) { + t.Parallel() + + body, err := BuildCodexModelsManifest([]string{"gpt-5.6-sol"}) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "low", models[0]["default_reasoning_level"]) + levels, ok := models[0]["supported_reasoning_levels"].([]any) + require.True(t, ok) + require.Len(t, levels, 6) + firstLevel, ok := levels[0].(map[string]any) + require.True(t, ok) + require.NotEqual(t, "none", firstLevel["effort"]) +} + +// Scenario: 专用图片生成模型不进入 Codex 主模型目录。 +func TestBuildCodexModelsManifestOmitsDedicatedImageModels(t *testing.T) { + t.Parallel() + + body, err := BuildCodexModelsManifest([]string{ + "grok-4.6", + "gpt-image-1", + "gpt-image-1.5", + "gpt-image-2", + "openai/gpt-image-2", + "gemini-2.5-flash-image", + "gemini-3.1-flash-image-preview", + "gemini-3-pro-image", + "google/gemini-3-pro-image", + "google/models/gemini-2.5-flash-image-preview", + "grok-imagine-image", + "grok-imagine-video", + "xai/grok-imagine-image-quality", + "grok-4.5", + }) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + slugs := make([]string, 0, len(models)) + for _, model := range models { + slug, _ := model["slug"].(string) + slugs = append(slugs, slug) + } + require.Equal(t, []string{"grok-4.6", "grok-4.5"}, slugs) +} + +func TestBuildCodexModelsManifestForGroupAdvertisesOfficialGrokResponsesImageInput(t *testing.T) { + t.Parallel() + + const groupID int64 = 701 + svc := &GatewayService{ + accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ + ID: 1, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token", + }, + }}, + }}, + } + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"grok-4.5"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.Equal(t, "Grok 4.5", models[0]["display_name"]) + require.Equal(t, []string{"low", "medium", "high"}, effortsFromManifestModel(t, models[0])) +} + +func TestBuildCodexModelsManifestForGroupAdvertisesOfficialOpenAIResponsesImageInput(t *testing.T) { + t.Parallel() + + const groupID int64 = 702 + svc := &GatewayService{ + accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ + ID: 2, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + }}, + }}, + } + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"gpt-5.6-sol"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) +} + +func TestBuildCodexModelsManifestForGroupUsesConservativeProviderImageCapabilities(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + model string + accounts []Account + modalities []any + }{ + { + name: "official Grok 4.6", + model: "grok-4.6", + accounts: []Account{{ + ID: 10, Platform: PlatformGrok, Type: AccountTypeOAuth, + }}, + modalities: []any{"text", "image"}, + }, + { + name: "official Grok Build vision host", + model: "grok-build-0.1", + accounts: []Account{{ + ID: 11, Platform: PlatformGrok, Type: AccountTypeOAuth, + }}, + modalities: []any{"text", "image"}, + }, + { + name: "official Grok 4.20 vision model", + model: "grok-4.20-0309-reasoning", + accounts: []Account{{ + ID: 23, Platform: PlatformGrok, Type: AccountTypeOAuth, + }}, + modalities: []any{"text", "image"}, + }, + { + name: "Grok 3 Mini is text only", + model: "grok-3-mini", + accounts: []Account{{ + ID: 24, Platform: PlatformGrok, Type: AccountTypeOAuth, + }}, + modalities: []any{"text"}, + }, + { + name: "Grok Composer has only Chat image bridge", + model: "grok-composer-2.5-fast", + accounts: []Account{{ + ID: 12, Platform: PlatformGrok, Type: AccountTypeOAuth, + }}, + modalities: []any{"text"}, + }, + { + name: "custom Grok host", + model: "grok-4.5", + accounts: []Account{{ + ID: 13, Platform: PlatformGrok, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://relay.example.test/v1"}, + }}, + modalities: []any{"text"}, + }, + { + name: "malformed Grok host", + model: "grok-4.5", + accounts: []Account{{ + ID: 19, Platform: PlatformGrok, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "::invalid::url"}, + }}, + modalities: []any{"text"}, + }, + { + name: "mixed official and custom Grok candidates", + model: "grok-4.5", + accounts: []Account{ + {ID: 14, Platform: PlatformGrok, Type: AccountTypeOAuth}, + { + ID: 15, Platform: PlatformGrok, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://relay.example.test/v1"}, + }, + }, + modalities: []any{"text"}, + }, + { + name: "DeepSeek V4", + model: "deepseek-v4-pro", + accounts: []Account{{ + ID: 16, Platform: PlatformDeepseek, Type: AccountTypeAPIKey, + }}, + modalities: []any{"text"}, + }, + { + name: "official OpenAI API key", + model: "gpt-5.6-sol", + accounts: []Account{{ + ID: 17, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + }}, + modalities: []any{"text", "image"}, + }, + { + name: "official OpenAI legacy text model", + model: "gpt-3.5-turbo", + accounts: []Account{{ + ID: 20, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + }}, + modalities: []any{"text"}, + }, + { + name: "custom OpenAI-compatible host", + model: "gpt-5.6-sol", + accounts: []Account{{ + ID: 18, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://openai-compatible.example.test/v1"}, + }}, + modalities: []any{"text"}, + }, + } + + for i, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + groupID := int64(710 + i) + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: tt.accounts, + }}} + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{tt.model}, + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, tt.modalities, models[0]["input_modalities"]) + }) + } +} + +func TestBuildCodexModelsManifestForGroupUsesExplicitCompositeResponsesRouteModel(t *testing.T) { + t.Parallel() + + const groupID int64 = 730 + routeRepo := compositeRouteRepoStub{routes: []CompositeModelRoute{{ + ID: 1, + GroupID: groupID, + PublicModel: "vision-alias", + MatchType: CompositeRouteMatchExact, + TargetPlatform: PlatformGrok, + UpstreamModel: "grok-4.5", + Endpoint: CompositeRouteEndpointResponses, + Enabled: true, + }}} + svc := &GatewayService{ + accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ID: 20, Platform: PlatformGrok, Type: AccountTypeOAuth}}, + }}, + compositeResolver: NewCompositeRouteResolver(routeRepo), + } + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"vision-alias"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) +} + +func TestBuildCodexModelsManifestForGroupUsesAccountMappingOwnershipAndMappedModel(t *testing.T) { + t.Parallel() + + const groupID int64 = 731 + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ + ID: 21, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "model_mapping": map[string]any{"vision-alias": "grok-4.5"}, + }, + }}, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"vision-alias"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) +} + +// Scenario: a Composite exact alias inherits metadata from its unique mapped target model. +func TestBuildCodexModelsManifestForGroupUsesMappedTargetMetadataForCompositeAlias(t *testing.T) { + t.Parallel() + + const groupID int64 = 733 + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ + ID: 23, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{"reasoning-alias": "claude-opus-4-8"}, + }, + }}, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"reasoning-alias"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "reasoning-alias", models[0]["slug"]) + require.Equal(t, "reasoning-alias", models[0]["display_name"]) + require.Equal(t, "Custom model routed through Sub2API.", models[0]["description"]) + require.Equal(t, []string{"low", "medium", "high", "xhigh", "max"}, effortsFromManifestModel(t, models[0])) +} + +// Scenario: conflicting targets on the same platform keep the public alias but do not guess capabilities. +func TestBuildCodexModelsManifestForGroupUsesSafeFallbackForConflictingAliasTargets(t *testing.T) { + t.Parallel() + + const groupID int64 = 734 + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: { + { + ID: 24, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{"shared-alias": "claude-opus-4-8"}, + }, + }, + { + ID: 25, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{"shared-alias": "claude-haiku-4-5-20251001"}, + }, + }, + }, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"shared-alias"}, + ) + require.NoError(t, err) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, "shared-alias", models[0]["slug"]) + require.Equal(t, "shared-alias", models[0]["display_name"]) + require.Equal(t, "Custom model routed through Sub2API.", models[0]["description"]) + require.Empty(t, effortsFromManifestModel(t, models[0])) +} + +// Scenario: a media-only target remains hidden even when exposed through an ordinary alias. +func TestBuildCodexModelsManifestForGroupOmitsDedicatedMediaTargetAlias(t *testing.T) { + t.Parallel() + + const groupID int64 = 735 + svc := &GatewayService{accountRepo: codexModelsVisibilityAccountRepo{byGroup: map[int64][]Account{ + groupID: {{ + ID: 26, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{"creative-alias": "gpt-image-2"}, + }, + }}, + }}} + + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"creative-alias"}, + ) + require.NoError(t, err) + require.Empty(t, decodeCodexManifestModels(t, body)) +} + +func TestBuildCodexModelsManifestForGroupLoadsAccountsOnce(t *testing.T) { + t.Parallel() + + const groupID int64 = 732 + repo := &countingCodexModelsAccountRepo{accounts: []Account{{ + ID: 22, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "vision-alias-a": "grok-4.5", + "vision-alias-b": "grok-4.6", + }, + }, + }}} + svc := &GatewayService{accountRepo: repo} + _, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: groupID, Platform: PlatformComposite}, + "", + []string{"vision-alias-a", "vision-alias-b", "deepseek-v4-pro"}, + ) + require.NoError(t, err) + require.Equal(t, int32(1), repo.calls.Load()) +} + +func TestBuildCodexModelsManifestForGroupUsesFallbackWhenTextOnlyPlatformHasNoSnapshot(t *testing.T) { + t.Parallel() + + repo := &countingCodexModelsAccountRepo{} + svc := &GatewayService{accountRepo: repo} + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: 733, Platform: PlatformDeepseek}, + "", + []string{"deepseek-v4-pro"}, + ) + require.NoError(t, err) + require.Equal(t, int32(1), repo.calls.Load()) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 1) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) +} + +func TestBuildCodexModelsManifestForGroupFallsBackWhenCapabilityLookupFails(t *testing.T) { + t.Parallel() + + repo := &countingCodexModelsAccountRepo{err: errors.New("account repository unavailable")} + svc := &GatewayService{accountRepo: repo} + body, err := svc.BuildCodexModelsManifestForGroup( + context.Background(), + &Group{ID: 734, Platform: PlatformComposite}, + "", + []string{"gpt-5.6-sol", "grok-4.5"}, + ) + require.NoError(t, err) + require.Equal(t, int32(1), repo.calls.Load()) + + models := decodeCodexManifestModels(t, body) + require.Len(t, models, 2) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.Equal(t, []any{"text"}, models[1]["input_modalities"]) +} + +func TestMergeGroupConfiguredCodexModelsInjectsCurrentGroupAliases(t *testing.T) { + t.Parallel() + + const groupID int64 = 71 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: { + { + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "deepseek-4-pro": "deepseek-v4-pro", + }, + }, + }, + }, + 72: { + { + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"other-group-model": "upstream-model"}, + }, + }, + }, + }, + }} + manifest := &CodexModelsManifest{ + Body: []byte(`{"models":[{"slug":"gpt-5.6","display_name":"GPT-5.6","unknown":{"kept":true}}],"metadata":{"version":1}}`), + } + + err := svc.MergeGroupConfiguredCodexModels( + context.Background(), + &Group{ID: groupID, Platform: PlatformOpenAI}, + manifest, + "", + ) + require.NoError(t, err) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 2) + require.Equal(t, "gpt-5.6", models[0]["slug"]) + require.Equal(t, map[string]any{"kept": true}, models[0]["unknown"]) + requireCompleteConfiguredCodexModel(t, models[1], "deepseek-4-pro") + require.EqualValues(t, 1_000_000, models[1]["context_window"]) + require.EqualValues(t, 1_000_000, models[1]["max_context_window"]) + require.Equal(t, "high", models[1]["default_reasoning_level"]) + require.Len(t, models[1]["supported_reasoning_levels"], 3) + require.NotContains(t, string(manifest.Body), "other-group-model") + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) +} + +// Scenario: OpenAI 分组存在账号模型配置时直接生成本地 Codex 清单。 +func TestBuildGroupConfiguredCodexModelsManifestUsesAdministratorConfiguration(t *testing.T) { + t.Parallel() + + const groupID int64 = 77 + reasoning := true + arkAccount := Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "glm-5.3": "glm-5.3", + "gpt-image-2": "gpt-image-2", + }, + }, + } + arkAccount.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "glm-5.3": { + ID: "glm-5.3", + DisplayName: "GLM 5.3", + Description: "Ark coding model", + Reasoning: &reasoning, + DefaultReasoningLevel: "medium", + SupportedReasoningLevels: []string{"low", "medium", "high"}, + InputModalities: []string{"text"}, + ContextWindow: 1_000_000, + }, + }}) + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: { + { + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + }, + arkAccount, + }, + }, + }} + group := &Group{ID: groupID, Platform: PlatformOpenAI} + + manifest, configured, err := svc.BuildGroupConfiguredCodexModelsManifest(context.Background(), group, "") + require.NoError(t, err) + require.True(t, configured) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + require.Equal(t, "glm-5.3", models[0]["slug"]) + require.Equal(t, "GLM 5.3", models[0]["display_name"]) + require.Equal(t, []string{"low", "medium", "high"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, "medium", models[0]["default_reasoning_level"]) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) + + notModified, configured, err := svc.BuildGroupConfiguredCodexModelsManifest( + context.Background(), + group, + "W/"+manifest.ETag, + ) + require.NoError(t, err) + require.True(t, configured) + require.True(t, notModified.NotModified) + require.Empty(t, notModified.Body) + require.Equal(t, manifest.ETag, notModified.ETag) +} + +// Scenario: OpenAI 通配映射展开组内精确选择,但不发布通配符 slug。 +func TestBuildGroupConfiguredCodexModelsManifestExpandsSelectedModelCoveredByWildcardMapping(t *testing.T) { + t.Parallel() + + const groupID int64 = 80 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: {{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-*": "gpt-5.6-sol"}, + }, + }}, + }, + }} + group := &Group{ + ID: groupID, + Platform: PlatformOpenAI, + ModelsListConfig: GroupModelsListConfig{ + Enabled: true, + Models: []string{"gpt-5.6"}, + }, + } + + manifest, configured, err := svc.BuildGroupConfiguredCodexModelsManifest(context.Background(), group, "") + require.NoError(t, err) + require.True(t, configured) + require.Equal(t, []string{"gpt-5.6"}, codexManifestModelSlugs(t, manifest.Body)) + require.NotContains(t, string(manifest.Body), "gpt-*") +} + +// Scenario: OpenAI 配置目录对暂时不可调度账号取能力交集,且不发布其独有模型。 +func TestBuildGroupConfiguredCodexModelsManifestIntersectsUnschedulableMappedAccounts(t *testing.T) { + t.Parallel() + + const groupID int64 = 79 + schedulable := newCodexCatalogMappedAccount( + 41, + "gpt-5.6-sol", + "GPT-5.6 Sol", + []string{"low", "medium", "high", "xhigh"}, + []string{"text", "image"}, + 1_000_000, + true, + nil, + ) + unschedulable := newCodexCatalogMappedAccount( + 42, + "glm-5.3", + "GLM 5.3", + []string{"low", "medium", "high"}, + []string{"text"}, + 272_000, + false, + map[string]any{"exclusive-model": "exclusive-upstream"}, + ) + svc := &OpenAIGatewayService{accountRepo: splitCodexModelsAccountRepo{ + schedulable: map[int64][]Account{groupID: {schedulable}}, + catalog: map[int64][]Account{groupID: {schedulable, unschedulable}}, + }} + + manifest, configured, err := svc.BuildGroupConfiguredCodexModelsManifest( + context.Background(), + &Group{ID: groupID, Platform: PlatformOpenAI}, + "", + ) + require.NoError(t, err) + require.True(t, configured) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + require.Equal(t, "my-coder", models[0]["slug"]) + require.Equal(t, "my-coder", models[0]["display_name"]) + require.Equal(t, []string{"low", "medium", "high"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.EqualValues(t, 272_000, models[0]["context_window"]) +} + +// Scenario: 没有管理员模型配置时保留现有上游发现路径。 +func TestBuildGroupConfiguredCodexModelsManifestFallsThroughWithoutConfiguration(t *testing.T) { + t.Parallel() + + const groupID int64 = 78 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: {{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}}, + }, + }} + + manifest, configured, err := svc.BuildGroupConfiguredCodexModelsManifest( + context.Background(), + &Group{ID: groupID, Platform: PlatformOpenAI}, + "", + ) + require.NoError(t, err) + require.False(t, configured) + require.Nil(t, manifest) +} + +func TestMergeGroupConfiguredCodexModelsFiltersAutoReviewByDefault(t *testing.T) { + t.Parallel() + + const groupID int64 = 74 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{}} + manifest := &CodexModelsManifest{ + Body: []byte(`{"models":[{"slug":"codex-auto-review","visibility":"list"},{"slug":"codex-auto-future","visibility":"list"},{"slug":"gpt-image-2","visibility":"list"},{"slug":"gpt-5.6","visibility":"list"}]}`), + } + + require.NoError(t, svc.MergeGroupConfiguredCodexModels( + context.Background(), + &Group{ID: groupID, Platform: PlatformOpenAI}, + manifest, + "", + )) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + require.Equal(t, "gpt-5.6", models[0]["slug"]) + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) +} + +// Scenario: OpenAI 账号映射不启用 Auto Review。 +func TestMergeGroupConfiguredCodexModelsFiltersAccountMappedAutoReviewByDefault(t *testing.T) { + t.Parallel() + + const groupID int64 = 75 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: { + { + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + openai.CodexUsageProbeModel: openai.CodexUsageProbeModel, + }, + }, + }, + }, + }, + }} + manifest := &CodexModelsManifest{ + Body: []byte(`{"models":[{"slug":"codex-auto-review","visibility":"hide","model_messages":{"auto_review":{"enabled":true}}},{"slug":"gpt-5.6","visibility":"list"}]}`), + } + + require.NoError(t, svc.MergeGroupConfiguredCodexModels( + context.Background(), + &Group{ID: groupID, Platform: PlatformOpenAI}, + manifest, + "", + )) + require.Equal(t, []string{"gpt-5.6"}, codexManifestModelSlugs(t, manifest.Body)) +} + +// Scenario: 启用的分组自定义列表允许 Auto Review。 +func TestMergeGroupConfiguredCodexModelsKeepsExplicitAutoReviewSelection(t *testing.T) { + t.Parallel() + + const groupID int64 = 76 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{}} + manifest := &CodexModelsManifest{ + Body: []byte(`{"models":[{"slug":"codex-auto-review","visibility":"list"},{"slug":"gpt-5.6","visibility":"list"}]}`), + } + group := &Group{ + ID: groupID, + Platform: PlatformOpenAI, + ModelsListConfig: GroupModelsListConfig{ + Enabled: true, + Models: []string{openai.CodexUsageProbeModel}, + }, + } + + require.NoError(t, svc.MergeGroupConfiguredCodexModels(context.Background(), group, manifest, "")) + require.Equal(t, []string{"codex-auto-review"}, codexManifestModelSlugs(t, manifest.Body)) +} + +func TestMergeGroupConfiguredCodexModelsHonorsCustomListAndFinalETag(t *testing.T) { + t.Parallel() + + const groupID int64 = 73 + svc := &OpenAIGatewayService{accountRepo: codexModelsVisibilityAccountRepo{ + byGroup: map[int64][]Account{ + groupID: { + { + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "deepseek-4-pro": "deepseek-v4-pro", + "hidden-alias": "hidden-upstream", + }, + }, + }, + }, + }, + }} + group := &Group{ + ID: groupID, + Platform: PlatformOpenAI, + ModelsListConfig: GroupModelsListConfig{ + Enabled: true, + Models: []string{"deepseek-4-pro"}, + }, + } + upstreamBody := []byte(`{"models":[{"slug":"gpt-5.6","display_name":"GPT-5.6"}]}`) + manifest := &CodexModelsManifest{Body: upstreamBody} + + require.NoError(t, svc.MergeGroupConfiguredCodexModels(context.Background(), group, manifest, "")) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + requireCompleteConfiguredCodexModel(t, models[0], "deepseek-4-pro") + + finalETag := manifest.ETag + second := &CodexModelsManifest{Body: upstreamBody} + require.NoError(t, svc.MergeGroupConfiguredCodexModels(context.Background(), group, second, finalETag)) + require.True(t, second.NotModified) + require.Empty(t, second.Body) + require.Equal(t, finalETag, second.ETag) +} + type codexModelsBlockingBody struct { ctx context.Context readStarted chan struct{} @@ -123,6 +1283,35 @@ func TestIsRetryableCodexModelsManifestTransportError(t *testing.T) { } } +func TestIsRetryableCodexModelsManifestStatus(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + statusCode int + useAPIKeyUpstream bool + retryable bool + }{ + {name: "api key 404", statusCode: http.StatusNotFound, useAPIKeyUpstream: true, retryable: true}, + {name: "api key 405", statusCode: http.StatusMethodNotAllowed, useAPIKeyUpstream: true, retryable: true}, + {name: "oauth 404", statusCode: http.StatusNotFound}, + {name: "oauth 405", statusCode: http.StatusMethodNotAllowed}, + {name: "api key 401", statusCode: http.StatusUnauthorized, useAPIKeyUpstream: true}, + {name: "oauth 401", statusCode: http.StatusUnauthorized, retryable: true}, + {name: "api key 400", statusCode: http.StatusBadRequest, useAPIKeyUpstream: true}, + {name: "api key 403", statusCode: http.StatusForbidden, useAPIKeyUpstream: true}, + {name: "rate limited", statusCode: http.StatusTooManyRequests, retryable: true}, + {name: "server error", statusCode: http.StatusServiceUnavailable, retryable: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.retryable, isRetryableCodexModelsManifestStatus(tt.statusCode, tt.useAPIKeyUpstream)) + }) + } +} + func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService { return &OpenAIGatewayService{ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{ @@ -407,7 +1596,7 @@ func TestFetchCodexModelsManifestMissingToken(t *testing.T) { } func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { - manifestBody := `{"models":[{"slug":"gpt-5.6"}]}` + manifestBody := `{"models":[{"slug":"deepseek-v4-pro"}]}` var gotRequest *http.Request var gotProxyURL string var gotAccountID int64 @@ -464,12 +1653,48 @@ func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { 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) - } + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + requireCompleteConfiguredCodexModel(t, models[0], "deepseek-v4-pro") + require.Equal(t, "DeepSeek V4 Pro", models[0]["display_name"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) + require.Equal(t, `W/"api-key-manifest"`, manifest.upstreamETag) +} + +// Scenario: 完整上游清单没有 ETag 时,最终正文仍生成强 ETag 并支持 304。 +func TestFetchCodexModelsManifestAPIKeyCompleteBodyWithoutUpstreamETagUsesFinalBodyETag(t *testing.T) { + completeBody, err := BuildCodexModelsManifest([]string{"custom-complete-model"}) + require.NoError(t, err) + + 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(bytes.NewReader(completeBody)), + }, nil + }} + svc := newCodexModelsAPIKeyTestService(upstream) + svc.accountRepo = codexModelsVisibilityAccountRepo{} + account := newCodexModelsAPIKeyTestAccount("https://upstream.example/v1") + group := &Group{ID: 82, Platform: PlatformOpenAI} + + first, err := svc.FetchCodexModelsManifest(context.Background(), account, "0.150.0", "") + require.NoError(t, err) + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(first, account)) + require.NoError(t, svc.MergeGroupConfiguredCodexModels(context.Background(), group, first, "")) + require.Equal(t, codexModelsManifestBodyETag(first.Body), first.ETag) + require.NotEmpty(t, first.ETag) + + second, err := svc.FetchCodexModelsManifest(context.Background(), account, "0.150.0", "") + require.NoError(t, err) + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(second, account)) + require.NoError(t, svc.MergeGroupConfiguredCodexModels(context.Background(), group, second, first.ETag)) + require.True(t, second.NotModified) + require.Empty(t, second.Body) + require.Equal(t, int32(1), calls.Load()) } func TestFetchCodexModelsManifestAPIKeyConvertsStandardOpenAIModelList(t *testing.T) { @@ -494,13 +1719,224 @@ func TestFetchCodexModelsManifestAPIKeyConvertsStandardOpenAIModelList(t *testin if err != nil { t.Fatalf("FetchCodexModelsManifest returned error: %v", err) } - if got, want := string(manifest.Body), `{"models":[{"slug":"gpt-5.6"},{"slug":"gpt-5.6-codex"}]}`; got != want { - t.Errorf("converted body: got %q, want %q", got, want) - } + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 2) + requireCompleteConfiguredCodexModel(t, models[0], "gpt-5.6") + requireCompleteConfiguredCodexModel(t, models[1], "gpt-5.6-codex") require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) require.Equal(t, `W/"openai-list"`, manifest.upstreamETag) } +func TestConvertOpenAIModelListToCodexManifestUsesCompleteDescriptors(t *testing.T) { + upstreamBody := `{"object":"list","data":[{"id":"gpt-5.5","object":"model"}]}` + + converted := convertOpenAIModelListToCodexManifest([]byte(upstreamBody)) + models := decodeCodexManifestModels(t, converted) + + require.Len(t, models, 1) + requireCompleteConfiguredCodexModel(t, models[0], "gpt-5.5") + require.Equal(t, "GPT-5.5", models[0]["display_name"]) + require.Equal(t, "medium", models[0]["default_reasoning_level"]) + require.Len(t, models[0]["supported_reasoning_levels"], 4) +} + +func TestCompleteAPIKeyCodexModelsManifestForClientPreservesProviderMetadata(t *testing.T) { + t.Parallel() + + svc := &OpenAIGatewayService{} + manifest := &CodexModelsManifest{ + Body: []byte(`{"models":[{"slug":"grok-4.6","description":"Provider supplied","model_messages":{"auto_review":{"enabled":true}},"truncation_policy":{"mode":"tokens"},"unknown":{"kept":true}}],"metadata":{"source":"upstream"}}`), + } + account := newCodexModelsAPIKeyTestAccount("https://upstream.example/v1") + + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(manifest, account)) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + requireCompleteConfiguredCodexModel(t, models[0], "grok-4.6") + require.Equal(t, "Provider supplied", models[0]["description"]) + require.Equal(t, map[string]any{"kept": true}, models[0]["unknown"]) + require.Equal(t, []any{"text"}, models[0]["input_modalities"]) + require.Equal(t, []string{"low", "medium", "high", "xhigh"}, effortsFromManifestModel(t, models[0])) + modelMessages, ok := models[0]["model_messages"].(map[string]any) + require.True(t, ok) + require.NotEmpty(t, modelMessages["instructions_template"]) + require.Equal(t, map[string]any{"enabled": true}, modelMessages["auto_review"]) + require.Equal(t, map[string]any{"mode": "tokens", "limit": float64(10_000)}, models[0]["truncation_policy"]) + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) + + var envelope map[string]any + require.NoError(t, json.Unmarshal(manifest.Body, &envelope)) + require.Equal(t, map[string]any{"source": "upstream"}, envelope["metadata"]) +} + +// Scenario: 标准 /models 型号列表优先使用已同步账号能力,再使用本地 descriptor 兜底。 +func TestCompleteAPIKeyCodexModelsManifestForClientUsesSyncedMetadataForConvertedModelList(t *testing.T) { + t.Parallel() + + reasoning := true + account := newCodexModelsAPIKeyTestAccount("https://upstream.example/v1") + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "future-reasoner": { + ID: "future-reasoner", DisplayName: "Future Reasoner", Description: "Synced upstream capability", + Reasoning: &reasoning, DefaultReasoningLevel: "ultra", + SupportedReasoningLevels: []string{"low", "high", "ultra"}, + InputModalities: []string{"text", "image"}, + ContextWindow: 999_000, + }, + }}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"object":"list","data":[{"id":"future-reasoner","object":"model"}]}`)), + }, nil + }} + svc := newCodexModelsAPIKeyTestService(upstream) + + manifest, err := svc.FetchCodexModelsManifest(context.Background(), account, "0.150.0", "") + require.NoError(t, err) + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(manifest, account)) + + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + require.Equal(t, "Future Reasoner", models[0]["display_name"]) + require.Equal(t, "Synced upstream capability", models[0]["description"]) + require.Equal(t, "ultra", models[0]["default_reasoning_level"]) + require.Equal(t, []string{"low", "high", "ultra"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.EqualValues(t, 999_000, models[0]["context_window"]) + require.EqualValues(t, 999_000, models[0]["max_context_window"]) +} + +func TestCompleteAPIKeyCodexModelsManifestForClientFillsMissingProviderFieldsWithoutOverwritingExplicitMetadata(t *testing.T) { + t.Parallel() + + reasoning := true + account := newCodexModelsAPIKeyTestAccount("https://upstream.example/v1") + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "provider-model": { + ID: "provider-model", DisplayName: "Synced Display", Description: "Synced description", + Reasoning: &reasoning, DefaultReasoningLevel: "ultra", + SupportedReasoningLevels: []string{"high", "ultra"}, + InputModalities: []string{"text", "image"}, + ContextWindow: 999_000, + }, + }}) + manifest := &CodexModelsManifest{Body: []byte(`{"models":[{ + "slug":"provider-model", + "description":"Provider supplied", + "context_window":64000, + "max_context_window":64000 + }]}`)} + svc := &OpenAIGatewayService{} + + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(manifest, account)) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + require.Equal(t, "Synced Display", models[0]["display_name"]) + require.Equal(t, "Provider supplied", models[0]["description"]) + require.Equal(t, "ultra", models[0]["default_reasoning_level"]) + require.Equal(t, []string{"high", "ultra"}, effortsFromManifestModel(t, models[0])) + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) + require.EqualValues(t, 64_000, models[0]["context_window"]) + require.EqualValues(t, 64_000, models[0]["max_context_window"]) +} + +// Scenario: 原生 manifest 的缺失字段在命中缓存后仍使用账号当前同步快照,而不是缓存中的本地默认值。 +func TestCompleteAPIKeyCodexModelsManifestForClientUsesCurrentSnapshotForCachedNativeManifest(t *testing.T) { + t.Parallel() + + reasoning := true + account := newCodexModelsAPIKeyTestAccount("https://upstream.example/v1") + setSnapshot := func(displayName string, contextWindow int64) { + account.SetUpstreamModelMetadataSnapshot(UpstreamModelMetadataSnapshot{Models: map[string]UpstreamModelMetadata{ + "deepseek-v4-pro": { + ID: "deepseek-v4-pro", DisplayName: displayName, + Reasoning: &reasoning, DefaultReasoningLevel: "ultra", + SupportedReasoningLevels: []string{"high", "ultra"}, + InputModalities: []string{"text", "image"}, + ContextWindow: contextWindow, + }, + }}) + } + setSnapshot("Synced DeepSeek", 256_000) + + 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: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"models":[{ + "slug":"deepseek-v4-pro", + "description":"Provider supplied" + }]}`)), + }, nil + }} + svc := newCodexModelsAPIKeyTestService(upstream) + + first, err := svc.FetchCodexModelsManifest(context.Background(), account, "0.150.0", "") + require.NoError(t, err) + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(first, account)) + firstModel := decodeCodexManifestModels(t, first.Body)[0] + require.Equal(t, "Synced DeepSeek", firstModel["display_name"]) + require.Equal(t, "Provider supplied", firstModel["description"]) + require.Equal(t, "ultra", firstModel["default_reasoning_level"]) + require.Equal(t, []string{"high", "ultra"}, effortsFromManifestModel(t, firstModel)) + require.Equal(t, []any{"text", "image"}, firstModel["input_modalities"]) + require.EqualValues(t, 256_000, firstModel["context_window"]) + + setSnapshot("Refreshed DeepSeek", 512_000) + second, err := svc.FetchCodexModelsManifest(context.Background(), account, "0.150.0", "") + require.NoError(t, err) + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(second, account)) + secondModel := decodeCodexManifestModels(t, second.Body)[0] + require.Equal(t, "Refreshed DeepSeek", secondModel["display_name"]) + require.Equal(t, "Provider supplied", secondModel["description"]) + require.EqualValues(t, 512_000, secondModel["context_window"]) + require.Equal(t, int32(1), calls.Load(), "second response should use the cached upstream source body") +} + +func TestCompleteAPIKeyCodexModelsManifestForClientMarksOnlyOfficialVisionGPTImageInput(t *testing.T) { + t.Parallel() + + svc := &OpenAIGatewayService{} + manifest := &CodexModelsManifest{Body: []byte(`{"models":[{"slug":"gpt-5.6-sol"},{"slug":"gpt-4o"},{"slug":"gpt-3.5-turbo"},{"slug":"gpt-4"}]}`)} + account := newCodexModelsAPIKeyTestAccount("") + + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(manifest, account)) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 4) + + bySlug := make(map[string]map[string]any, len(models)) + for _, model := range models { + slug, ok := model["slug"].(string) + require.True(t, ok) + bySlug[slug] = model + } + for _, slug := range []string{"gpt-5.6-sol", "gpt-4o"} { + require.Equal(t, []any{"text", "image"}, bySlug[slug]["input_modalities"]) + require.Equal(t, true, bySlug[slug]["supports_image_detail_original"]) + } + for _, slug := range []string{"gpt-3.5-turbo", "gpt-4"} { + require.Equal(t, []any{"text"}, bySlug[slug]["input_modalities"]) + require.Equal(t, false, bySlug[slug]["supports_image_detail_original"]) + } + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) +} + +func TestCompleteAPIKeyCodexModelsManifestForClientFiltersOfficialNonAgentModels(t *testing.T) { + t.Parallel() + + svc := &OpenAIGatewayService{} + manifest := &CodexModelsManifest{Body: []byte(`{"models":[{"slug":"gpt-5.6-sol"},{"slug":"gpt-4o-realtime-preview"},{"slug":"gpt-4o-mini-tts"},{"slug":"text-embedding-3-large"},{"slug":"omni-moderation-latest"},{"slug":"o4-mini"},{"slug":"codex-mini-latest"}]}`)} + account := newCodexModelsAPIKeyTestAccount("") + + require.NoError(t, svc.CompleteAPIKeyCodexModelsManifestForClient(manifest, account)) + require.Equal(t, []string{"gpt-5.6-sol", "o4-mini", "codex-mini-latest"}, codexManifestModelSlugs(t, manifest.Body)) + require.Equal(t, codexModelsManifestBodyETag(manifest.Body), manifest.ETag) +} + func TestAdjustAPIKeyCodexModelsManifest(t *testing.T) { tests := []struct { name string @@ -572,17 +2008,27 @@ func TestFetchCodexModelsManifestOAuthPreservesResponsesLite(t *testing.T) { require.Equal(t, manifestBody, string(manifest.Body)) } -func TestConvertOpenAIModelListToCodexManifest(t *testing.T) { +func TestConvertOpenAIModelListToCompleteCodexManifest(t *testing.T) { + t.Parallel() + + body := []byte(`{"object":"list","data":[{"id":"deepseek-v4-flash"},{"id":"deepseek-v4-pro"}]}`) + models := decodeCodexManifestModels(t, convertOpenAIModelListToCodexManifest(body)) + + require.Len(t, models, 2) + requireCompleteConfiguredCodexModel(t, models[0], "deepseek-v4-flash") + requireCompleteConfiguredCodexModel(t, models[1], "deepseek-v4-pro") + require.Equal(t, "DeepSeek V4 Flash", models[0]["display_name"]) + require.Equal(t, "DeepSeek V4 Pro", models[1]["display_name"]) + require.EqualValues(t, 1_000_000, models[0]["context_window"]) + require.EqualValues(t, 1_000_000, models[1]["context_window"]) +} + +func TestConvertOpenAIModelListToCodexManifestLeavesUnsupportedBodiesUnchanged(t *testing.T) { tests := []struct { name string body string want string }{ - { - name: "standard list", - body: `{"object":"list","data":[{"id":"m-1"},{"id":"m-2"}]}`, - want: `{"models":[{"slug":"m-1"},{"slug":"m-2"}]}`, - }, { name: "codex manifest unchanged", body: `{"models":[{"slug":"m-1"}]}`, @@ -910,6 +2356,77 @@ func TestFetchCodexModelsManifestAPIKeyFreshCacheHandlesETagLocally(t *testing.T } } +func TestFetchCodexModelsManifestAPIKeyCacheSurvivesClientMutation(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(`{"object":"list","data":[{"id":"model-a"},{"id":"model-b"}]}`)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + s.accountRepo = codexModelsVisibilityAccountRepo{} + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + + first, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + require.NoError(t, err) + require.Contains(t, string(first.Body), "model-a") + require.Contains(t, string(first.Body), "model-b") + + require.NoError(t, s.CompleteAPIKeyCodexModelsManifestForClient(first, account)) + require.NoError(t, s.MergeGroupConfiguredCodexModels( + context.Background(), + &Group{ + ID: 81, + Platform: PlatformOpenAI, + ModelsListConfig: GroupModelsListConfig{ + Enabled: true, + Models: []string{"model-a"}, + }, + }, + first, + "", + )) + require.Equal(t, []string{"model-a"}, codexManifestModelSlugs(t, first.Body)) + + second, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + require.NoError(t, err) + require.Equal(t, []string{"model-a", "model-b"}, codexManifestModelSlugs(t, second.Body)) + require.Equal(t, int32(1), calls.Load()) + + var wg sync.WaitGroup + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + manifest, fetchErr := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + require.NoError(t, fetchErr) + require.NoError(t, s.MergeGroupConfiguredCodexModels( + context.Background(), + &Group{ + ID: 82, + Platform: PlatformOpenAI, + ModelsListConfig: GroupModelsListConfig{ + Enabled: true, + Models: []string{"model-b"}, + }, + }, + manifest, + "", + )) + require.Equal(t, []string{"model-b"}, codexManifestModelSlugs(t, manifest.Body)) + }() + } + wg.Wait() + + third, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + require.NoError(t, err) + require.Equal(t, []string{"model-a", "model-b"}, codexManifestModelSlugs(t, third.Body)) + require.Equal(t, int32(1), calls.Load()) +} + func TestFetchCodexModelsManifestAPIKeyCacheKeyIsolatesRequestIdentity(t *testing.T) { var calls atomic.Int32 upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { @@ -966,7 +2483,10 @@ func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing. 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") { + switch { + case strings.Contains(req.URL.Host, "large-source"): + body = `{"object":"list","data":[{"id":"model-a","padding":"` + strings.Repeat("x", 1<<20) + `"}]}` + case strings.Contains(req.URL.Host, "large"): body = `{"models":[],"padding":"` + strings.Repeat("x", (1<<20)+1) + `"}` } return &http.Response{ @@ -990,8 +2510,12 @@ func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing. 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) + largeSource := newCodexModelsAPIKeyTestAccount("https://large-source.example") + largeSource.ID = 4 + fetch(largeSource) + fetch(largeSource) + if got := calls.Load(); got != 5 { + t.Fatalf("body-size bounded cache calls: got %d, want 5", got) } for i := int64(10); i < 75; i++ { @@ -1002,14 +2526,14 @@ func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing. 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) + if got := calls.Load(); got != 70 { + t.Fatalf("most recent cache entry was not retained: calls=%d, want 70", 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) + if got := calls.Load(); got != 71 { + t.Errorf("oldest cache entry was not evicted: calls=%d, want 71", got) } } @@ -1462,7 +2986,7 @@ func TestFetchCodexModelsManifestAPIKeyUpstreamError(t *testing.T) { } } -func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T) { +func TestFetchCodexModelsManifestAPIKeyUsesOfficialOpenAIModelsEndpoint(t *testing.T) { tests := []struct { name string baseURL string @@ -1474,23 +2998,32 @@ func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T 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 + var gotURL string + s := newCodexModelsAPIKeyTestService(&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(`{"object":"list","data":[{"id":"gpt-5.6-sol"}]}`)), + }, nil }}) - _, err := s.FetchCodexModelsManifest( + manifest, 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)) - } + require.NoError(t, err) + parsedURL, parseErr := url.Parse(gotURL) + require.NoError(t, parseErr) + require.Equal(t, "api.openai.com", strings.ToLower(parsedURL.Hostname())) + require.Equal(t, "/v1/models", parsedURL.Path) + require.Equal(t, "0.144.0", parsedURL.Query().Get("client_version")) + models := decodeCodexManifestModels(t, manifest.Body) + require.Len(t, models, 1) + requireCompleteConfiguredCodexModel(t, models[0], "gpt-5.6-sol") + require.Equal(t, []any{"text", "image"}, models[0]["input_modalities"]) }) } } diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go index b0e98e5926..a22288b70a 100644 --- a/backend/internal/service/openai_gateway_cc_pipeline.go +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -186,6 +186,10 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest( if err != nil { return nil, fmt.Errorf("build upstream request: %w", err) } + // 记录本次实际选择的协议端点,供错误日志和用量日志在没有 + // OpenAIForwardResult(例如 503/传输失败)时使用。每次发送都覆盖, + // 避免 Gin context 在账号 failover 尝试之间残留旧端点。 + SetActualOpenAIUpstreamEndpoint(c, "/v1/chat/completions") upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) upstreamReq.Header.Set("Content-Type", "application/json") upstreamReq.Header.Set("Authorization", "Bearer "+bearerToken) diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 69c1458fdc..a3a992175c 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -60,6 +60,10 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( defaultMappedModel string, ) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + ClearActualOpenAIUpstreamEndpoint(c) + if shouldForwardOpenAIResponsesViaRawChatCompletions(account) { + SetActualOpenAIUpstreamEndpoint(c, "/v1/chat/completions") + } setCodexToolNameReverse(c, nil) if _, err := s.prepareCodexAccountIdentitySource(ctx, c, account); err != nil { return nil, err diff --git a/backend/internal/service/openai_gateway_cn_fixes_test.go b/backend/internal/service/openai_gateway_cn_fixes_test.go index f5cd28ee39..576f8bfd42 100644 --- a/backend/internal/service/openai_gateway_cn_fixes_test.go +++ b/backend/internal/service/openai_gateway_cn_fixes_test.go @@ -11,6 +11,7 @@ package service import ( "context" + "errors" "net/http" "net/http/httptest" "testing" @@ -142,3 +143,112 @@ func TestHandle403_CNProviderStructured403TempUnschedulableFirstHit(t *testing.T require.Equal(t, 1, repo.tempCalls) require.Contains(t, repo.lastTempReason, "(1/3)") } + +func TestIsCNProviderConcurrencyLimit403_ExactClassification(t *testing.T) { + kimi := &Account{Platform: PlatformKimi} + + require.True(t, isCNProviderConcurrencyLimit403(kimi, kimiConcurrentRequestLimitMessage)) + require.True(t, isCNProviderConcurrencyLimit403(kimi, " "+kimiConcurrentRequestLimitMessage+"\n")) + + for name, tc := range map[string]struct { + account *Account + message string + }{ + "permission denied": {kimi, "You do not have permission to access this resource."}, + "generic concurrency wording": {kimi, "concurrent request limit reached"}, + "near match missing punctuation": {kimi, "You've reached your concurrent request limit. Please wait for your ongoing requests to finish and try again"}, + "other CN provider": {&Account{Platform: PlatformZhipu}, kimiConcurrentRequestLimitMessage}, + "non CN provider": {&Account{Platform: PlatformOpenAI}, kimiConcurrentRequestLimitMessage}, + "nil account": {nil, kimiConcurrentRequestLimitMessage}, + } { + t.Run(name, func(t *testing.T) { + require.False(t, isCNProviderConcurrencyLimit403(tc.account, tc.message)) + }) + } +} + +func TestHandle403_OtherCNProviderWithKimiConcurrencyMessageUsesNormalPolicy(t *testing.T) { + repo := &rateLimitAccountRepoStub{} + counter := &openAI403CounterCacheStub{counts: []int64{openAI403DisableThreshold}} + blocker := &runtimeBlockRecorder{} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + service.SetOpenAI403CounterCache(counter) + service.SetAccountRuntimeBlocker(blocker) + account := &Account{ID: 405, Platform: PlatformZhipu, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), account, http.StatusForbidden, http.Header{}, + []byte(`{"error":{"message":"You've reached your concurrent request limit. Please wait for your ongoing requests to finish and try again."}}`), + ) + + require.True(t, shouldDisable) + require.Equal(t, 1, repo.setErrorCalls, "non-Kimi CN provider must retain the normal permanent-error policy") + require.Equal(t, 0, repo.tempCalls) + require.Empty(t, counter.counts, "normal CN 403 policy must consume the counter result") + require.Equal(t, []string{"auth_error"}, blocker.reasons, "the Kimi-specific runtime block must not apply") +} + +func TestHandle403_CNProviderConcurrencyLimitAlwaysUsesTemporaryCooldown(t *testing.T) { + repo := &rateLimitAccountRepoStub{} + counter := &openAI403CounterCacheStub{counts: []int64{openAI403DisableThreshold}} + blocker := &runtimeBlockRecorder{} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + service.SetOpenAI403CounterCache(counter) + service.SetAccountRuntimeBlocker(blocker) + account := &Account{ID: 403, Platform: PlatformKimi, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), account, http.StatusForbidden, http.Header{}, + []byte(`{"error":{"message":"You've reached your concurrent request limit. Please wait for your ongoing requests to finish and try again."}}`), + ) + + require.True(t, shouldDisable, "the request must still fail over to another account") + require.Equal(t, 0, repo.setErrorCalls) + require.Equal(t, 1, repo.tempCalls) + require.Contains(t, repo.lastTempReason, cnConcurrencyLimitReasonPrefix) + require.Equal(t, []int64{openAI403DisableThreshold}, counter.counts, "transient concurrency 403 must bypass the permanent-error counter") + require.Len(t, blocker.accounts, 1) + require.Equal(t, cnConcurrencyLimitReasonPrefix, blocker.reasons[0]) + require.True(t, blocker.until[0].After(time.Now())) +} + +func TestHandle403_KimiConcurrencyLimitRepositoryFailureKeepsRuntimeBlock(t *testing.T) { + repo := &rateLimitAccountRepoStub{tempErr: errors.New("repository unavailable")} + counter := &openAI403CounterCacheStub{counts: []int64{openAI403DisableThreshold}} + blocker := &runtimeBlockRecorder{} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + service.SetOpenAI403CounterCache(counter) + service.SetAccountRuntimeBlocker(blocker) + account := &Account{ID: 406, Platform: PlatformKimi, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), account, http.StatusForbidden, http.Header{}, + []byte(`{"error":{"message":"You've reached your concurrent request limit. Please wait for your ongoing requests to finish and try again."}}`), + ) + + require.True(t, shouldDisable, "the current request must fail over even when persistence fails") + require.Equal(t, 1, repo.tempCalls, "the temporary cooldown should still be persisted when possible") + require.Equal(t, 0, repo.setErrorCalls, "persistence failure must not fall back to permanent account error") + require.Equal(t, []int64{openAI403DisableThreshold}, counter.counts, "persistence failure must not enter the permanent-error counter path") + require.Len(t, blocker.accounts, 1, "the in-memory runtime block must survive repository failure") + require.Same(t, account, blocker.accounts[0]) + require.Equal(t, cnConcurrencyLimitReasonPrefix, blocker.reasons[0]) + require.True(t, blocker.until[0].After(time.Now())) +} + +func TestHandle403_CNProviderNearMatchRetainsNormalPermanentErrorPolicy(t *testing.T) { + repo := &rateLimitAccountRepoStub{} + counter := &openAI403CounterCacheStub{counts: []int64{openAI403DisableThreshold}} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + service.SetOpenAI403CounterCache(counter) + account := &Account{ID: 404, Platform: PlatformKimi, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), account, http.StatusForbidden, http.Header{}, + []byte(`{"error":{"message":"You've reached your concurrent request limit. Please contact support."}}`), + ) + + require.True(t, shouldDisable) + require.Equal(t, 1, repo.setErrorCalls, "non-exact 403 must retain existing permission/auth protection") + require.Equal(t, 0, repo.tempCalls) +} diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index ebf12675c9..304af59909 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -20,6 +20,15 @@ import ( // Forward forwards request to OpenAI API func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + ClearActualOpenAIUpstreamEndpoint(c) + if shouldForwardOpenAIResponsesViaRawChatCompletions(account) { + SetActualOpenAIUpstreamEndpoint(c, "/v1/chat/completions") + } + filteredBody, filterErr := filterOpenAIResponsesNoneReasoningEffortForAccount(account, body) + if filterErr != nil { + return nil, filterErr + } + body = filteredBody clearGrokResponsesClientToolMapping(c) clearOpenAIResponsesClientToolMapping(c) clearOpenAIResponsesNamespaceNames(c) @@ -112,7 +121,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } if shouldStripOpenAIResponsesInputNamespaces(account, wsDecision.Transport, passthroughEnabled) { keepToolCallNamespaces := shouldKeepOpenAIResponsesToolCallNamespaces( - account, wsDecision.Transport, passthroughEnabled, compactPath, + account, wsDecision.Transport, passthroughEnabled, compactPath, body, ) body, err = stripOpenAIResponsesInputNamespaces(body, keepToolCallNamespaces) if err != nil { @@ -474,13 +483,29 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if decodeErr != nil { return nil, decodeErr } + // Responses OAuth 与 Chat 兼容入口保持一致:纯文本 system 可以无损提升后删除, + // JSON object 模式仍需在 input 中保留 JSON 指令供上游兼容校验。 + omitPromotedSystemMessages := !strings.EqualFold( + strings.TrimSpace(gjson.GetBytes(body, "text.format.type").String()), + "json_object", + ) codexResult := codexTransformResult{} if compatMessagesBridge { - codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{IsCodexCLI: isCodexCLI, IsCompact: isCompactRequest, SkipDefaultInstructions: true, PreserveToolCallIDs: true}) + codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{ + IsCodexCLI: isCodexCLI, + IsCompact: isCompactRequest, + SkipDefaultInstructions: true, + PreserveToolCallIDs: true, + OmitPromotedSystemMessagesFromInput: omitPromotedSystemMessages, + }) ensureCodexOAuthInstructionsField(decoded) markDecodedModified() } else { - codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest) + codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{ + IsCodexCLI: isCodexCLI, + IsCompact: isCompactRequest, + OmitPromotedSystemMessagesFromInput: omitPromotedSystemMessages, + }) } if codexResult.Error != nil { c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": codexResult.Error.Error()}}) diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 17c4d633d1..a494be8adf 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -34,6 +34,10 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( defaultMappedModel string, ) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + ClearActualOpenAIUpstreamEndpoint(c) + if shouldForwardOpenAIResponsesViaRawChatCompletions(account) { + SetActualOpenAIUpstreamEndpoint(c, "/v1/chat/completions") + } setCodexToolNameReverse(c, nil) if _, err := s.prepareCodexAccountIdentitySource(ctx, c, account); err != nil { return nil, err diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index b87ce9df87..e3f19f335e 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -1678,7 +1678,14 @@ func (s *OpenAIGatewayService) newOpenAIStreamFailoverError( }, }) retryableOnSameAccount := openAIStreamFailedEventRetryableOnSameAccount(account, payload, message) - failoverErr := s.newOpenAIAccountFailoverError(account, statusCode, headers, payload, message, shouldDisable, retryableOnSameAccount) + // 流终止事件承载在 HTTP 200 内,外层响应头描述的是成功流状态,而不是语义上的 + // 429 事件。仅在配额分类时忽略这些头;故障转移错误仍保留它们,使 Retry-After + // 和请求 ID 能继续传递给后续处理。 + classificationHeaders := headers + if statusCode == http.StatusTooManyRequests { + classificationHeaders = nil + } + failoverErr := s.newOpenAIAccountFailoverErrorWithClassificationHeaders(account, statusCode, headers, classificationHeaders, payload, message, shouldDisable, retryableOnSameAccount) if failoverErr.IsCredentialFailure() || failoverErr.RequestScopedTransient { return failoverErr } diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 23cc01d670..dec535c543 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -55,6 +55,69 @@ func buildOpenAIResponsesURLForPlatform(platform string, base string) string { return buildOpenAIResponsesURL(base) } +func shouldPreserveOpenAIResponsesNoneReasoningEffort(account *Account) bool { + if account == nil { + return false + } + if account.IsOpenAIOAuthLike() { + return true + } + if !account.IsOpenAIApiKey() { + return false + } + baseURL := strings.TrimSpace(account.GetCredential("base_url")) + return baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) +} + +// Codex 0.149.0 needs a single advertised effort to directly select a visible +// non-reasoning model. Treat that catalog-only "none" value as omission for +// compatible upstreams, while preserving official OpenAI request semantics. +func filterOpenAIResponsesNoneReasoningEffortForAccount(account *Account, body []byte) ([]byte, error) { + if len(body) == 0 || shouldPreserveOpenAIResponsesNoneReasoningEffort(account) { + return body, nil + } + + out := body + for _, path := range []string{"reasoning.effort", "reasoning_effort"} { + effort := gjson.GetBytes(out, path) + if effort.Type != gjson.String || !strings.EqualFold(strings.TrimSpace(effort.String()), "none") { + continue + } + next, err := sjson.DeleteBytes(out, path) + if err != nil { + return body, fmt.Errorf("strip %s none placeholder: %w", path, err) + } + out = next + } + if reasoning := gjson.GetBytes(out, "reasoning"); reasoning.IsObject() && len(reasoning.Map()) == 0 { + next, err := sjson.DeleteBytes(out, "reasoning") + if err != nil { + return body, fmt.Errorf("strip empty reasoning object: %w", err) + } + out = next + } + return out, nil +} + +func deleteOpenAIResponsesNoneReasoningEffortFromObject(account *Account, body map[string]any) { + if body == nil || shouldPreserveOpenAIResponsesNoneReasoningEffort(account) { + return + } + if effort, ok := body["reasoning_effort"].(string); ok && strings.EqualFold(strings.TrimSpace(effort), "none") { + delete(body, "reasoning_effort") + } + reasoning, ok := body["reasoning"].(map[string]any) + if !ok { + return + } + if effort, ok := reasoning["effort"].(string); ok && strings.EqualFold(strings.TrimSpace(effort), "none") { + delete(reasoning, "effort") + } + if len(reasoning) == 0 { + delete(body, "reasoning") + } +} + // normalizeDeepSeekResponsesRequestBody 适配 DeepSeek 无状态 Responses 端点: // 强制 store=false 并清除 previous_response_id(官方 /responses 不支持服务端 // 状态存储,携带这些字段会被拒绝)。非 deepseek responses 协议账号原样返回。 diff --git a/backend/internal/service/openai_gateway_request_body_reasoning_test.go b/backend/internal/service/openai_gateway_request_body_reasoning_test.go index cbe24815cd..4805eeeb24 100644 --- a/backend/internal/service/openai_gateway_request_body_reasoning_test.go +++ b/backend/internal/service/openai_gateway_request_body_reasoning_test.go @@ -271,6 +271,66 @@ func TestNormalizeOpenAIParallelToolCallsWithoutTools(t *testing.T) { require.False(t, gjson.GetBytes(normalized, "parallel_tool_calls").Exists()) } +func TestFilterOpenAIResponsesNoneReasoningEffortForAccount(t *testing.T) { + tests := []struct { + name string + account *Account + body string + wantNested bool + wantFlat bool + wantSummary bool + wantReasoning bool + }{ + { + name: "custom compatible endpoint strips none placeholders", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "https://compat.example/v1"}}, + body: `{"reasoning":{"effort":"none"},"reasoning_effort":"NONE"}`, + wantReasoning: false, + }, + { + name: "third-party platform keeps other reasoning members", + account: &Account{Platform: PlatformGrok, Type: AccountTypeAPIKey}, + body: `{"reasoning":{"effort":" none ","summary":"auto"}}`, + wantSummary: true, + wantReasoning: true, + }, + { + name: "non-none effort is unchanged", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "https://compat.example/v1"}}, + body: `{"reasoning":{"effort":"high"},"reasoning_effort":"low"}`, + wantNested: true, + wantFlat: true, + wantReasoning: true, + }, + { + name: "official OpenAI API key preserves none", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + body: `{"reasoning":{"effort":"none"},"reasoning_effort":"none"}`, + wantNested: true, + wantFlat: true, + wantReasoning: true, + }, + { + name: "OpenAI OAuth preserves none", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + body: `{"reasoning":{"effort":"none"}}`, + wantNested: true, + wantReasoning: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := filterOpenAIResponsesNoneReasoningEffortForAccount(tt.account, []byte(tt.body)) + require.NoError(t, err) + require.Equal(t, tt.wantNested, gjson.GetBytes(got, "reasoning.effort").Exists()) + require.Equal(t, tt.wantFlat, gjson.GetBytes(got, "reasoning_effort").Exists()) + require.Equal(t, tt.wantSummary, gjson.GetBytes(got, "reasoning.summary").Exists()) + require.Equal(t, tt.wantReasoning, gjson.GetBytes(got, "reasoning").Exists()) + }) + } +} + // Lite 工具迁移到 input[].additional_tools 后,仍应按有工具请求处理。 func TestNormalizeOpenAIParallelToolCallsWithoutTools_KeepsResponsesLiteAdditionalTools(t *testing.T) { liteBody := []byte(`{"input":[{"type":"message","role":"user","content":"hi"},{"type":"additional_tools","tools":[{"type":"function","name":"spawn_agent"}]}],"parallel_tool_calls":false}`) diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback_test.go b/backend/internal/service/openai_gateway_responses_chat_fallback_test.go index 750c45e8f2..bb994b1aa4 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback_test.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback_test.go @@ -39,11 +39,13 @@ func TestForwardResponses_ForceChatCompletionsRoutesNonStreamingToChatCompletion cfg: rawChatCompletionsTestConfig(), httpUpstream: upstream, } + SetActualOpenAIUpstreamEndpoint(c, "/v1/responses") result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body) require.NoError(t, err) require.NotNil(t, result) require.Equal(t, "http://upstream.example/v1/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, "/v1/chat/completions", GetActualOpenAIUpstreamEndpoint(c)) require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(upstream.lastReq.Context())) require.Equal(t, "hello", gjson.GetBytes(upstream.lastBody, "messages.0.content").String()) require.False(t, gjson.GetBytes(upstream.lastBody, "input").Exists()) @@ -55,6 +57,36 @@ func TestForwardResponses_ForceChatCompletionsRoutesNonStreamingToChatCompletion require.False(t, result.Stream) } +// Scenario: 第三方无推理模型不收到兼容档位。 +func TestForwardResponses_ForceChatCompletionsOmitsNoneReasoningEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"company-coding-model","input":"hello","reasoning":{"effort":"none"},"stream":false}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"chatcmpl_none","object":"chat.completion","model":"company-coding-model","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`, + )), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + } + + result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "company-coding-model", gjson.GetBytes(upstream.lastBody, "model").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "reasoning_effort").Exists()) + require.Nil(t, result.ReasoningEffort) +} + func TestForwardResponses_PassthroughFlagWithUnsupportedResponsesUsesAccountMapping(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 2a4a4f3d09..a90bad3e25 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -28,6 +28,7 @@ const ( ) var explicitOpenAIHeaderSessionNames = []string{ + "session-id", "session_id", "conversation_id", openCodeSessionAffinityHeader, @@ -145,7 +146,7 @@ func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body // GenerateSessionHash generates a sticky-session hash for OpenAI requests. // // Priority: -// 1. Header: session_id +// 1. Header: session-id / session_id // 2. Header: conversation_id // 3. Header: x-session-affinity / x-session-id / x-opencode-session (OpenCode) // 4. Header: x-conversation-id (CodeBuddy) @@ -1173,6 +1174,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } // ============ Layer 1: Sticky session ============ + // A healthy sticky account whose bounded wait queue is full may be used as a + // one-request capacity spillover in Layer 2. Keep that spillover temporary: + // rewriting the durable binding here would make a short burst migrate the + // whole conversation to a cache-cold account. + stickySpillover := false if sessionHash != "" { accountID := stickyAccountID if accountID > 0 && !isExcluded(accountID) { @@ -1214,6 +1220,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex MaxWaiting: cfg.StickySessionMaxWaiting, }) } + stickySpillover = true } } } @@ -1365,7 +1372,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex if selectErr != nil { return nil, true, selectErr } - if sessionHash != "" && !gatewayProfitControlGateActive(ctx) { + if sessionHash != "" && !stickySpillover && !gatewayProfitControlGateActive(ctx) { _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) } return selection, true, nil @@ -1404,7 +1411,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex if selectErr != nil { return nil, selectErr } - if sessionHash != "" && !gatewayProfitControlGateActive(ctx) { + if sessionHash != "" && !stickySpillover && !gatewayProfitControlGateActive(ctx) { _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) } return selection, nil diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index c4d085b024..5ce98fba0f 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -319,6 +319,16 @@ func SetActualOpenAIUpstreamEndpoint(c *gin.Context, endpoint string) { } } +// ClearActualOpenAIUpstreamEndpoint 清理当前转发尝试记录的端点。 +// Handler 会在账号 failover 尝试间复用同一个 Gin context,因此每次尝试 +// 都必须从无残留状态开始。 +func ClearActualOpenAIUpstreamEndpoint(c *gin.Context) { + if c == nil { + return + } + c.Set(openAIUpstreamEndpointContextKey, "") +} + // GetActualOpenAIUpstreamEndpoint returns the endpoint recorded by the latest // forwarding attempt in this request. func GetActualOpenAIUpstreamEndpoint(c *gin.Context) string { diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 1e41298cc1..568c948627 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -391,6 +391,7 @@ func TestOpenAIGatewayService_ClientSessionHeaderPriority(t *testing.T) { name string value string }{ + {name: "session-id", value: "codex-session"}, {name: "session_id", value: "generic-session"}, {name: "conversation_id", value: "generic-conversation"}, {name: openCodeSessionAffinityHeader, value: "opencode-affinity"}, @@ -416,6 +417,31 @@ func TestOpenAIGatewayService_ClientSessionHeaderPriority(t *testing.T) { require.Equal(t, "body-session", svc.ExtractSessionID(c, body)) } +func TestOpenAIGatewayService_CodexSessionIDKeepsReconnectHashStable(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + c.Request.Header.Set("session-id", "codex-reconnect-session") + + svc := &OpenAIGatewayService{} + warmup := []byte(`{ + "type":"response.create", + "model":"gpt-5.6-sol", + "generate":false, + "tools":[{"type":"custom","name":"exec"}], + "input":[{"role":"user","content":"warmup"}] + }`) + business := []byte(`{ + "type":"response.create", + "model":"gpt-5.6-sol", + "input":[{"role":"user","content":"install codex"}] + }`) + + require.Equal(t, svc.GenerateSessionHash(c, warmup), svc.GenerateSessionHash(c, business)) + require.Equal(t, "codex-reconnect-session", svc.ExtractSessionID(c, business)) +} + func TestOpenAIGatewayService_ClientSessionHeadersIgnorePerRequestIDs(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() @@ -1181,6 +1207,49 @@ func TestOpenAISelectAccountWithLoadAwareness_StickyWaitPlan(t *testing.T) { } } +func TestOpenAISelectAccountWithLoadAwareness_StickyCapacitySpilloverKeepsBinding(t *testing.T) { + sessionHash := "sticky-spillover" + groupID := int64(1) + repo := stubOpenAIAccountRepo{ + accounts: []Account{ + {ID: 1, Platform: PlatformOpenAI, Status: StatusActive, Schedulable: true, Concurrency: 6, Priority: 1, GroupIDs: []int64{groupID}}, + {ID: 2, Platform: PlatformOpenAI, Status: StatusActive, Schedulable: true, Concurrency: 6, Priority: 1, GroupIDs: []int64{groupID}}, + }, + } + cache := &stubGatewayCache{ + sessionBindings: map[string]int64{"openai:" + sessionHash: 1}, + } + concurrencyCache := stubConcurrencyCache{ + acquireResults: map[int64]bool{1: false, 2: true}, + waitCounts: map[int64]int{1: 1}, + loadMap: map[int64]*AccountLoadInfo{ + 1: {AccountID: 1, LoadRate: 100}, + 2: {AccountID: 2, LoadRate: 10}, + }, + } + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + cfg.Gateway.Scheduling.StickySessionMaxWaiting = 1 + + svc := &OpenAIGatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, err := svc.SelectAccountWithLoadAwareness(context.Background(), &groupID, sessionHash, "gpt-4", nil) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(2), selection.Account.ID, "capacity spillover should use the other account for this request") + require.True(t, selection.Acquired) + require.Equal(t, int64(1), cache.sessionBindings["openai:"+sessionHash], "capacity spillover must not migrate the durable sticky binding") + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + func TestOpenAISelectAccountWithLoadAwareness_PrefersLowerLoad(t *testing.T) { groupID := int64(1) repo := stubOpenAIAccountRepo{ diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 4a4e9ab7fc..70c000e321 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -333,7 +333,20 @@ func (s *OpenAIGatewayService) newOpenAIAccountFailoverError( shouldDisable bool, retryableOnSameAccount bool, ) *UpstreamFailoverError { - oauth429Retry := s.shouldRetryOpenAIOAuth429OnSameAccount(account, statusCode, shouldDisable) + return s.newOpenAIAccountFailoverErrorWithClassificationHeaders(account, statusCode, responseHeaders, responseHeaders, responseBody, upstreamMsg, shouldDisable, retryableOnSameAccount) +} + +func (s *OpenAIGatewayService) newOpenAIAccountFailoverErrorWithClassificationHeaders( + account *Account, + statusCode int, + responseHeaders http.Header, + classificationHeaders http.Header, + responseBody []byte, + upstreamMsg string, + shouldDisable bool, + retryableOnSameAccount bool, +) *UpstreamFailoverError { + oauth429Retry := s.shouldRetryOpenAIOAuth429OnSameAccountWithResponse(account, statusCode, shouldDisable, classificationHeaders, responseBody) failoverErr := newOpenAIUpstreamFailoverError( statusCode, responseHeaders, diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 9374a6f7c7..ed11746cdb 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -37,6 +37,13 @@ type OpenAIImagesUpstreamError struct { Message string Param string UpstreamRequestID string + + // SynthesizedFromModelText marks an error the gateway inferred from the + // model's plain-text output instead of reading it off a structured upstream + // error frame. Such a verdict describes this one turn ("the model answered + // with words instead of an image"), not the account — see + // shouldCoolOpenAIImagesToolForError. + SynthesizedFromModelText bool } func (e *OpenAIImagesUpstreamError) Error() string { @@ -328,6 +335,26 @@ func openAIImageUploadToDataURL(upload OpenAIImagesUpload) (string, error) { return "data:" + contentType + ";base64," + base64.StdEncoding.EncodeToString(upload.Data), nil } +// openAIImagesSelfBuiltRequestContextKey marks a request whose upstream body was +// fully constructed by buildOpenAIImagesResponsesRequest, i.e. tool_choice and the +// matching image_generation tool are always both present and never client-controlled. +type openAIImagesSelfBuiltRequestContextKey struct{} + +func withOpenAIImagesSelfBuiltRequest(ctx context.Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, openAIImagesSelfBuiltRequestContextKey{}, true) +} + +func isOpenAIImagesSelfBuiltRequest(ctx context.Context) bool { + if ctx == nil { + return false + } + selfBuilt, _ := ctx.Value(openAIImagesSelfBuiltRequestContextKey{}).(bool) + return selfBuilt +} + func buildOpenAIImagesResponsesRequest(parsed *OpenAIImagesRequest, toolModel string) ([]byte, error) { if parsed == nil { return nil, fmt.Errorf("parsed images request is required") @@ -711,6 +738,10 @@ func openAIImagesTextFallbackErrorForText(text string) *OpenAIImagesUpstreamErro ErrorType: "upstream_error", Code: "image_generation_unavailable", Message: "Upstream did not execute image generation", + // Inferred from the model's own words, not from an upstream error frame: + // good enough to fail this turn over to another account, not evidence that + // this account's image tool is down for the next 30 minutes. + SynthesizedFromModelText: true, } } @@ -1775,6 +1806,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( if err != nil { return nil, err } + upstreamCtx = withOpenAIImagesSelfBuiltRequest(upstreamCtx) upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, true, parsed.StickySessionSeed(), false) if err != nil { return nil, err @@ -1922,6 +1954,26 @@ const ( openAIImagesOAuthUnavailableReason = "openai_images_oauth_tool_unavailable" ) +// shouldCoolOpenAIImagesToolForError decides whether an image_generation_unavailable +// verdict is durable enough to park the account's image tool for +// openAIImagesOAuthUnavailableCooldown. +// +// Only an upstream error frame that names the condition qualifies. A verdict the +// gateway synthesized from the model's plain-text reply does not: it merely says +// this prompt produced words instead of an image, which is prompt-dependent and +// happens on healthy accounts. Writing a 30-minute account-level cooldown from it +// is doubly wrong because the very same error is classified retryable +// (IsOpenAIImagesRetryableUpstreamError: status >= 500) and drives +// newOpenAIAccountFailoverError — so one such reply walks the pool and cools every +// account the retry touches. +// +// This mirrors the rule the alpha/search path already states in words: a +// tool-endpoint failure "仍允许本次请求换号,但不修改任何账号状态" +// (see shouldApplyOpenAIAlphaSearchAccountErrorSideEffects). +func shouldCoolOpenAIImagesToolForError(upstreamErr *OpenAIImagesUpstreamError) bool { + return upstreamErr != nil && !upstreamErr.SynthesizedFromModelText +} + func (s *OpenAIGatewayService) coolOpenAIImagesOAuthTool(ctx context.Context, account *Account) { if s == nil || s.accountRepo == nil || account == nil || account.Platform != PlatformOpenAI { return @@ -2017,7 +2069,9 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthResponseError( responseBody := openAIImagesUpstreamErrorResponseBody(upstreamErr) if upstreamErr.Code == "image_generation_unavailable" { - s.coolOpenAIImagesOAuthTool(ctx, account) + if shouldCoolOpenAIImagesToolForError(upstreamErr) { + s.coolOpenAIImagesOAuthTool(ctx, account) + } if responseWritten { return err } diff --git a/backend/internal/service/openai_images_tool_cooldown_test.go b/backend/internal/service/openai_images_tool_cooldown_test.go new file mode 100644 index 0000000000..0f38c200d7 --- /dev/null +++ b/backend/internal/service/openai_images_tool_cooldown_test.go @@ -0,0 +1,177 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// issue #6171:v0.1.181 起,/v1/images/generations 只要上游"回文字没回图",账号就被 +// 写 30 分钟 openai:image_generation 模型级冷却。该判据是**请求级**的(这个 prompt +// 这一轮模型选择了说话),却被当成**账号级**能力失效;又因为同一个错误被判为 +// 可重试(502)并驱动 failover,一次闲聊回复会沿着号池逐个把账号冷却掉。 + +// countingModelRateLimitRepo 记录 SetModelRateLimit 调用,用于断言"没写账号状态"。 +type countingModelRateLimitRepo struct { + accountRepoStub + calls int + scopes []string +} + +func (r *countingModelRateLimitRepo) SetModelRateLimit(_ context.Context, _ int64, scope string, _ time.Time, _ ...string) error { + r.calls++ + r.scopes = append(r.scopes, scope) + return nil +} + +func newImagesCooldownContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + return c, rec +} + +func imagesCooldownAccount() *Account { + return &Account{ID: 77, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "img-oauth"} +} + +func TestShouldCoolOpenAIImagesToolForError(t *testing.T) { + cases := []struct { + name string + err *OpenAIImagesUpstreamError + want bool + }{ + { + name: "nil_error", + err: nil, + want: false, + }, + { + // 网关从模型文字里推断出来的判据:只说明这一轮没出图。 + name: "synthesized_from_model_text", + err: &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadGateway, + Code: "image_generation_unavailable", + SynthesizedFromModelText: true, + }, + want: false, + }, + { + // 上游自己在 error 帧里点名该状态:这才是账号级证据,保持冷却。 + name: "structured_upstream_error_frame", + err: &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadGateway, + Code: "image_generation_unavailable", + }, + want: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.want, shouldCoolOpenAIImagesToolForError(tc.err)) + }) + } +} + +// 主复现:文字兜底判据不得写账号级冷却。 +func TestHandleOpenAIImagesOAuthResponseError_TextFallbackDoesNotCoolAccount(t *testing.T) { + c, _ := newImagesCooldownContext(t) + repo := &countingModelRateLimitRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := imagesCooldownAccount() + + upstreamErr := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.") + require.NotNil(t, upstreamErr) + require.Equal(t, "image_generation_unavailable", upstreamErr.Code) + + err := svc.handleOpenAIImagesOAuthResponseError( + context.Background(), c, account, "gpt-image-2", "https://upstream.example/v1/responses", + &http.Response{StatusCode: http.StatusOK, Header: http.Header{}}, + OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), upstreamErr, + ) + + require.Zero(t, repo.calls, "模型闲聊不构成账号级证据,不得写 30 分钟冷却") + + // 换号行为必须原样保留:本 PR 只撤销账号状态写入,不动 failover。 + var failover *UpstreamFailoverError + require.True(t, errors.As(err, &failover), "仍应触发换号,got %T", err) +} + +// 对照不变式:上游 error 帧点名该状态时仍然冷却,否则等于把功能整个废掉。 +func TestHandleOpenAIImagesOAuthResponseError_StructuredUnavailableStillCoolsAccount(t *testing.T) { + c, _ := newImagesCooldownContext(t) + repo := &countingModelRateLimitRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := imagesCooldownAccount() + + upstreamErr := &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadGateway, + ErrorType: "upstream_error", + Code: "image_generation_unavailable", + Message: "image generation tool is not available for this account", + } + + _ = svc.handleOpenAIImagesOAuthResponseError( + context.Background(), c, account, "gpt-image-2", "https://upstream.example/v1/responses", + &http.Response{StatusCode: http.StatusOK, Header: http.Header{}}, + OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), upstreamErr, + ) + + require.Equal(t, 1, repo.calls, "结构化上游证据仍须写冷却") + require.Equal(t, []string{openAIImageGenerationRateLimitKey}, repo.scopes) +} + +// 标记必须打在文字兜底的两个入口上,且不影响违规拦截分支的判定。 +func TestOpenAIImagesTextFallback_MarksSynthesizedVerdicts(t *testing.T) { + t.Run("plain_text_reply_is_synthesized", func(t *testing.T) { + err := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.") + require.NotNil(t, err) + require.True(t, err.SynthesizedFromModelText) + require.Equal(t, "image_generation_unavailable", err.Code) + require.Equal(t, http.StatusBadGateway, err.StatusCode) + }) + + t.Run("body_entrypoint_is_synthesized", func(t *testing.T) { + body := []byte("event: response.completed\n" + + `data: {"type":"response.completed","response":{"id":"r","status":"completed",` + + `"output":[{"type":"message","content":[{"type":"output_text","text":"I drafted a prompt for you."}]}]}}` + + "\n\n") + err := openAIImagesTextFallbackError(body) + require.NotNil(t, err) + require.True(t, err.SynthesizedFromModelText) + }) + + t.Run("content_policy_branch_unchanged", func(t *testing.T) { + err := openAIImagesTextFallbackErrorForText("Blocked by our content policy.") + require.NotNil(t, err) + require.Equal(t, "content_policy_violation", err.Code) + require.Equal(t, http.StatusBadRequest, err.StatusCode) + // 该分支本来就不走冷却(Code 不匹配),标记与否都不改变行为; + // 断言它没有被顺手打标,避免语义漂移。 + require.False(t, err.SynthesizedFromModelText) + }) + + t.Run("empty_text_yields_no_error", func(t *testing.T) { + require.Nil(t, openAIImagesTextFallbackErrorForText(" ")) + }) +} + +// 级联的前提条件:该错误确实是可重试的,所以会带着"已写冷却"的副作用换号。 +// 这条用例把前提钉死,避免以后有人把 502 改成非重试后误以为本修复多余。 +func TestOpenAIImagesTextFallback_RemainsRetryableAndThusCascades(t *testing.T) { + err := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.") + require.NotNil(t, err) + require.True(t, IsOpenAIImagesRetryableUpstreamError(err), + "文字兜底判据是可重试的——正因如此,写账号冷却会沿号池级联") +} diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index fb9a6fee27..fa750fa5b5 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -132,6 +132,44 @@ func TestOpenAIGatewayService_ResponsesUnknownModelDoesNotFallbackToGPT54(t *tes require.True(t, rec.Code >= http.StatusBadRequest) } +func TestOpenAIGatewayService_OAuthResponsesPromotesSystemMessageWithoutDuplication(t *testing.T) { + gin.SetMode(gin.TestMode) + + const systemPrompt = "Unique system prefix for Responses token accounting." + const existingInstructions = "Existing instructions." + body := []byte(`{"model":"gpt-5.4","stream":false,"instructions":"` + existingInstructions + `","input":[{"role":"system","content":"` + systemPrompt + `"},{"role":"user","content":"hello"}]}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstream := &httpUpstreamRecorder{err: errors.New("stop after capture")} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 124, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + Status: StatusActive, + Schedulable: true, + } + + result, err := svc.Forward(context.Background(), c, account, body) + + require.Error(t, err) + require.Nil(t, result) + require.NotEmpty(t, upstream.lastBody) + require.Equal(t, systemPrompt+"\n\n"+existingInstructions, gjson.GetBytes(upstream.lastBody, "instructions").String()) + require.Equal(t, int64(1), gjson.GetBytes(upstream.lastBody, "input.#").Int()) + require.Equal(t, "user", gjson.GetBytes(upstream.lastBody, "input.0.role").String()) + require.Equal(t, 1, strings.Count(string(upstream.lastBody), systemPrompt)) +} + func TestOpenAIGatewayService_NativeResponsesBodyModificationPreservesHTMLChars(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_responses_namespace.go b/backend/internal/service/openai_responses_namespace.go index f4fd47b13b..fdb0212402 100644 --- a/backend/internal/service/openai_responses_namespace.go +++ b/backend/internal/service/openai_responses_namespace.go @@ -77,25 +77,48 @@ func shouldStripOpenAIResponsesInputNamespaces(account *Account, transport OpenA // 故 OAuth 非 compact 请求必须保留。 // - compact 端点的 schema 不含该字段,携带即 400 `Unknown parameter: // input[N].namespace`(issue #4761 正文),故 compact 一律清理。 -// - API Key 出口是标准 Responses API(api.openai.com 或自定义 base_url),同样 -// 不认识该字段,维持全量清理;否则只能退化成 -// openai_responses_rejected_field_retry 的逐项删除,6 次上限根本盖不住长历史。 +// - API Key 出口默认按标准 Responses API 处理并清理该字段;但当请求本身声明 +// namespace 工具时,上游显然使用了 namespace 扩展,此时必须保留调用项上的 +// namespace,否则声明与历史调用会失配并触发 Missing namespace。 // - 摊平模式下调用项已被改写成平名,残留 namespace 指向的声明已不存在,一律清理。 func shouldKeepOpenAIResponsesToolCallNamespaces( account *Account, transport OpenAIUpstreamTransport, passthroughEnabled bool, compactPath bool, + body []byte, ) bool { - if account == nil || !account.IsOpenAIOAuthLike() { + if account == nil { return false } if compactPath { return false } + if account.IsOpenAIApiKey() { + return hasOpenAIResponsesNamespaceToolDeclaration(body) + } + if !account.IsOpenAIOAuthLike() { + return false + } return !shouldFlattenOpenAIResponsesNamespaces(account, transport, passthroughEnabled, compactPath) } +func hasOpenAIResponsesNamespaceToolDeclaration(body []byte) bool { + tools := gjson.GetBytes(body, "tools") + if !tools.IsArray() { + return false + } + found := false + tools.ForEach(func(_, tool gjson.Result) bool { + if strings.EqualFold(strings.TrimSpace(tool.Get("type").String()), "namespace") { + found = true + return false + } + return true + }) + return found +} + // openAIResponsesToolCallItemTypes 是携带 namespace 的调用项类型集合。与 // removeOpenAIResponsesRejectedNamespaceAtIndex 的反应式白名单保持一致;codex-rs // protocol/src/models.rs 中只有 FunctionCall 与 CustomToolCall 序列化 namespace, diff --git a/backend/internal/service/openai_responses_namespace_forward_test.go b/backend/internal/service/openai_responses_namespace_forward_test.go index e99d0199fd..9ef4f370b4 100644 --- a/backend/internal/service/openai_responses_namespace_forward_test.go +++ b/backend/internal/service/openai_responses_namespace_forward_test.go @@ -66,6 +66,29 @@ func TestOpenAIGatewayService_OAuthPreservesCodexNamespaceTools(t *testing.T) { require.Empty(t, openAIResponsesNamespaceNames(c)) } +// API Key 自定义上游若接受 namespace 工具声明,也要求历史 function_call 原样携带 +// namespace。声明仍为命名空间工具却清掉调用项字段,会触发 Missing namespace。 +func TestOpenAIGatewayService_APIKeyPreservesDeclaredNamespaceToolCalls(t *testing.T) { + body := []byte(codexNamespaceRequestBody) + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + newOpenAIRejectedFieldTestResponse(http.StatusOK, namespaceForwardOKResponse), + }} + c := newOpenAIRejectedFieldTestContext(body) + + result, err := newOpenAIRejectedFieldTestService(upstream).Forward( + context.Background(), c, newOpenAIRejectedFieldTestAccount(), body, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 1) + forwarded := upstream.bodies[0] + + require.True(t, gjson.GetBytes(forwarded, `tools.#(type=="namespace")`).Exists()) + require.Equal(t, "collaboration", gjson.GetBytes(forwarded, "input.0.namespace").String()) + require.False(t, gjson.GetBytes(forwarded, "input.1.namespace").Exists()) +} + // compact 端点 schema 更窄:input[].namespace 会 400 Unknown parameter(issue #4761), // 且没有证据表明它接受 namespace 工具声明。compact 只做历史摘要、不需要模型寻址工具, // 因此保持既有的摊平 + 全量清理行为,不随默认值翻转扩大风险面。 diff --git a/backend/internal/service/openai_responses_namespace_test.go b/backend/internal/service/openai_responses_namespace_test.go index 7a9b47b7f5..f2ed2d8901 100644 --- a/backend/internal/service/openai_responses_namespace_test.go +++ b/backend/internal/service/openai_responses_namespace_test.go @@ -78,6 +78,7 @@ func TestShouldKeepOpenAIResponsesToolCallNamespaces(t *testing.T) { transport OpenAIUpstreamTransport passthroughEnabled bool compactPath bool + body []byte want bool }{ // 上游按 namespace 解析历史调用,缺字段会 400 "Missing namespace for function_call"。 @@ -92,15 +93,20 @@ func TestShouldKeepOpenAIResponsesToolCallNamespaces(t *testing.T) { // WSv2 + compact 是唯一「不摊平但仍必须清理」的组合,钉住 compact 判定本身, // 使其不会被误当成可由 shouldFlatten 推导出的冗余分支。 {name: "oauth_compact_wsv2_strips", account: oauth, transport: OpenAIUpstreamTransportResponsesWebsocketV2, compactPath: true, want: false}, - // API Key 出口是标准 Responses API,不认识该字段。 - {name: "apikey_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, want: false}, + // API Key 默认按标准 Responses API 清理;请求显式声明 namespace 工具时, + // 自定义上游需要原样接收对应的历史调用。 + {name: "apikey_without_namespace_tool_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, want: false}, + {name: "apikey_with_namespace_tool_keeps", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":"namespace","name":"mcp__codex_app","tools":[]}]}`), want: true}, + {name: "apikey_with_mixed_case_namespace_tool_keeps", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":" Namespace ","name":"mcp__codex_app","tools":[]}]}`), want: true}, + {name: "apikey_function_tool_with_namespace_field_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":"function","name":"automation_update","namespace":"mcp__codex_app"}]}`), want: false}, + {name: "apikey_compact_with_namespace_tool_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, compactPath: true, body: []byte(`{"tools":[{"type":"namespace","name":"mcp__codex_app","tools":[]}]}`), want: false}, {name: "setup_token_keeps", account: setupToken, transport: OpenAIUpstreamTransportHTTPSSE, want: true}, {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, shouldKeepOpenAIResponsesToolCallNamespaces( - tt.account, tt.transport, tt.passthroughEnabled, tt.compactPath, + tt.account, tt.transport, tt.passthroughEnabled, tt.compactPath, tt.body, )) }) } diff --git a/backend/internal/service/openai_ws_forwarder_logutil.go b/backend/internal/service/openai_ws_forwarder_logutil.go index ac5a3720b1..ddb37d8380 100644 --- a/backend/internal/service/openai_ws_forwarder_logutil.go +++ b/backend/internal/service/openai_ws_forwarder_logutil.go @@ -65,7 +65,10 @@ func resolveOpenAIWSSessionHeaders(c *gin.Context, promptCacheKey string) openAI ConversationSource: "none", } if c != nil && c.Request != nil { - if sessionID := strings.TrimSpace(c.Request.Header.Get("session_id")); sessionID != "" { + if sessionID := strings.TrimSpace(c.Request.Header.Get("session-id")); sessionID != "" { + resolution.SessionID = sessionID + resolution.SessionSource = "header_session-id" + } else if sessionID := strings.TrimSpace(c.Request.Header.Get("session_id")); sessionID != "" { resolution.SessionID = sessionID resolution.SessionSource = "header_session_id" } diff --git a/backend/internal/service/openai_ws_forwarder_logutil_test.go b/backend/internal/service/openai_ws_forwarder_logutil_test.go new file mode 100644 index 0000000000..eaf8b5f22b --- /dev/null +++ b/backend/internal/service/openai_ws_forwarder_logutil_test.go @@ -0,0 +1,37 @@ +package service + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestResolveOpenAIWSSessionHeadersPrefersCodexHyphenHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + c.Request.Header.Set("session-id", "codex-session") + c.Request.Header.Set("session_id", "legacy-session") + + resolution := resolveOpenAIWSSessionHeaders(c, "prompt-cache") + + require.Equal(t, "codex-session", resolution.SessionID) + require.Equal(t, "header_session-id", resolution.SessionSource) +} + +func TestResolveOpenAIWSSessionHeadersFallsBackToLegacyHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + c.Request.Header.Set("session_id", "legacy-session") + + resolution := resolveOpenAIWSSessionHeaders(c, "prompt-cache") + + require.Equal(t, "legacy-session", resolution.SessionID) + require.Equal(t, "header_session_id", resolution.SessionSource) +} diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index e775104fa2..6923352d7d 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -804,6 +804,73 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T require.Equal(t, isolateOpenAIUpstreamSessionID(0, account, "conv-oauth-1"), captureDialer.lastHeaders.Get("conversation_id")) } +func TestOpenAIGatewayService_Forward_WSv2_OAuthSanitizesInvalidNativeToolItemID(t *testing.T) { + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + groupID := int64(5662) + c.Set("api_key", &APIKey{GroupID: &groupID}) + + cfg := newOpenAIWSV2TestConfig() + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + 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_oauth_tool_history","model":"gpt-5.6-sol","usage":{"input_tokens":1,"output_tokens":1}}}`), + }} + captureDialer := &openAIWSCaptureDialer{conn: captureConn} + pool := newOpenAIWSConnPool(cfg) + pool.setClientDialerForTest(captureDialer) + + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: &httpUpstreamRecorder{}, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + openaiWSPool: pool, + } + + account := &Account{ + ID: 5662, + Name: "openai-oauth-ws-tool-history", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{"access_token": "test-oauth-token"}, + Extra: map[string]any{ + "responses_websockets_v2_enabled": true, + }, + } + + body := []byte(`{"model":"gpt-5.6-sol","stream":false,"instructions":"Continue the task.","input":[{"type":"custom_tool_call","id":"fc_hotfix_probe","call_id":"fc_hotfix","name":"exec","input":"pwd","status":"completed"},{"type":"custom_tool_call_output","call_id":"fc_hotfix","output":"done"}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + + captureConn.mu.Lock() + requestPayload := cloneMapStringAny(captureConn.lastWrite) + captureConn.mu.Unlock() + requestJSON := requestToJSONString(requestPayload) + require.Equal(t, "response.create", gjson.Get(requestJSON, "type").String()) + require.False(t, gjson.Get(requestJSON, "input.0.id").Exists(), "stale fc_* ID must not be replayed as a native custom_tool_call ID") + require.Equal(t, "ctc_hotfix", gjson.Get(requestJSON, "input.0.call_id").String()) + require.Equal(t, "custom_tool_call_output", gjson.Get(requestJSON, "input.1.type").String()) + require.Equal(t, "ctc_hotfix", gjson.Get(requestJSON, "input.1.call_id").String()) +} + func TestOpenAIGatewayService_Forward_WSv2_OAuthOriginatorCompatibility(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 2bb0c200b5..a6c0528045 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -115,7 +115,7 @@ func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(account *Account, payloa return threshold > 0 && int64(payloadBytes) >= threshold } -func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) { +func prepareOpenAIWSHTTPBridgeBody(account *Account, payload []byte) ([]byte, error) { var body map[string]any if err := decodeOpenAIJSONUseNumber(payload, &body); err != nil { return nil, err @@ -126,6 +126,7 @@ func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) { delete(body, "type") delete(body, "generate") delete(body, "previous_response_id") + deleteOpenAIResponsesNoneReasoningEffortFromObject(account, body) body["stream"] = true return json.Marshal(body) } @@ -305,7 +306,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( } responseModelObserver := &upstreamResponseModelObserver{} - body, err := prepareOpenAIWSHTTPBridgeBody(payload) + body, err := prepareOpenAIWSHTTPBridgeBody(account, payload) if err != nil { return nil, fmt.Errorf("prepare http bridge body: %w", err) } @@ -836,7 +837,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( } func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, seedPayload, currentPayload []byte, originalModel string) (string, error) { - body, err := prepareOpenAIWSHTTPBridgeBody(seedPayload) + body, err := prepareOpenAIWSHTTPBridgeBody(account, seedPayload) if err != nil { return "", err } diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 13752a55ef..bbff773d76 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -31,7 +31,7 @@ func TestResolveOpenAIWSClientFirstMessageTimeout(t *testing.T) { } func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) { - body, err := prepareOpenAIWSHTTPBridgeBody([]byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":false,"previous_response_id":"resp_prev","input":"hi","sequence":900719925474099312345}`)) + body, err := prepareOpenAIWSHTTPBridgeBody(nil, []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":false,"previous_response_id":"resp_prev","input":"hi","sequence":900719925474099312345}`)) require.NoError(t, err) require.False(t, gjson.GetBytes(body, "type").Exists()) require.False(t, gjson.GetBytes(body, "generate").Exists()) @@ -40,10 +40,26 @@ func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) { require.True(t, gjson.GetBytes(body, "stream").Bool()) require.Equal(t, "hi", gjson.GetBytes(body, "input").String()) require.Equal(t, "900719925474099312345", gjson.GetBytes(body, "sequence").Raw) - _, err = prepareOpenAIWSHTTPBridgeBody([]byte(`{"type":"response.create"}{"trailing":true}`)) + _, err = prepareOpenAIWSHTTPBridgeBody(nil, []byte(`{"type":"response.create"}{"trailing":true}`)) require.Error(t, err) } +func TestPrepareOpenAIWSHTTPBridgeBodyStripsNoneReasoningForCompatibleEndpoint(t *testing.T) { + payload := []byte(`{"type":"response.create","model":"company-coding-model","reasoning":{"effort":"none"},"input":"hi"}`) + compatible := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{ + "base_url": "https://compat.example/v1", + }} + + body, err := prepareOpenAIWSHTTPBridgeBody(compatible, payload) + require.NoError(t, err) + require.False(t, gjson.GetBytes(body, "reasoning.effort").Exists()) + require.False(t, gjson.GetBytes(body, "reasoning").Exists()) + + officialBody, err := prepareOpenAIWSHTTPBridgeBody(&Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, payload) + require.NoError(t, err) + require.Equal(t, "none", gjson.GetBytes(officialBody, "reasoning.effort").String()) +} + func TestProxyOpenAIWSHTTPBridgeTurn_UpstreamDefaultServiceTierWinsOverRequest(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/ratelimit_cn_providers.go b/backend/internal/service/ratelimit_cn_providers.go index 4aab5a8751..17239e4bbc 100644 --- a/backend/internal/service/ratelimit_cn_providers.go +++ b/backend/internal/service/ratelimit_cn_providers.go @@ -26,6 +26,33 @@ const cnBalanceExtraSuffixLow = "balance_low" // 其他子系统(阈值/限流/401)写入的临时停调。 const cnBalanceLowReasonPrefix = "cn_balance_low" +const kimiConcurrentRequestLimitMessage = "You've reached your concurrent request limit. Please wait for your ongoing requests to finish and try again." + +const cnConcurrencyLimitReasonPrefix = "cn_concurrency_limit" + +func isCNProviderConcurrencyLimit403(account *Account, upstreamMsg string) bool { + return account != nil && account.Platform == PlatformKimi && + strings.TrimSpace(upstreamMsg) == kimiConcurrentRequestLimitMessage +} + +func (s *RateLimitService) handleCNProviderConcurrencyLimit403( + ctx context.Context, + account *Account, +) { + until := time.Now().Add(time.Duration(openAI403CooldownMinutesDefault) * time.Minute) + reason := cnConcurrencyLimitReasonPrefix + ": " + kimiConcurrentRequestLimitMessage + s.notifyAccountSchedulingBlocked(account, until, cnConcurrencyLimitReasonPrefix) + if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, until, reason); err != nil { + slog.Warn("cn_concurrency_limit_set_temp_unschedulable_failed", "account_id", account.ID, "error", err) + return + } + slog.Info("cn_provider_concurrency_limited", + "account_id", account.ID, + "platform", account.Platform, + "until", until.UTC(), + ) +} + // cnBalanceLowReason 构造余额不足临时停调的 reason(带稳定前缀)。 func cnBalanceLowReason(upstreamMsg string) string { if upstreamMsg = strings.TrimSpace(upstreamMsg); upstreamMsg != "" { diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 0a52925cd9..4d6be56b06 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -75,6 +75,8 @@ const ( const ( openAIImageRateLimitDefaultCooldown = time.Minute openAIImageRateLimitReason = "openai_image_rate_limited" + openAIImageCapabilityLossCooldown = 30 * time.Minute + openAIImageCapabilityLossReason = "openai_image_capability_lost" ) var openAIImageTryAgainPattern = regexp.MustCompile(`(?i)try again in\s+([0-9]+(?:\.[0-9]+)?)\s*(ms|s|sec|secs|second|seconds|m|min|mins|minute|minutes)`) @@ -936,6 +938,13 @@ func (s *RateLimitService) handle403(ctx context.Context, account *Account, upst if account.Platform == PlatformAntigravity { return s.handleAntigravity403(ctx, account, upstreamMsg, responseBody) } + // Kimi reports its transient per-account concurrency/business limit as a 403. + // Keep the normal 403 failover signal (true), but never feed this exact message + // into the escalating 403 counter that can permanently mark the account error. + if isCNProviderConcurrencyLimit403(account, upstreamMsg) { + s.handleCNProviderConcurrencyLimit403(ctx, account) + return true + } // 国产供应商与 openai 同口径:HTML 403(CDN/代理拦截页)不构成账号失效证据, // 且 403 在 failover 状态集里会被逐账号重放——直接 SetError 会让一个坏请求/ // 一层坏代理连环永久禁用整组账号。走 HTML 豁免 + N 次累计 + 临时冷却。 @@ -2183,6 +2192,44 @@ func (s *RateLimitService) HandleOpenAIImageRateLimit(ctx context.Context, accou return true } +func (s *RateLimitService) HandleOpenAIImageCapabilityLoss(ctx context.Context, account *Account, statusCode int, responseBody []byte) bool { + if s == nil || account == nil || s.accountRepo == nil { + return false + } + if account.Platform != PlatformOpenAI { + return false + } + if !account.ShouldHandleErrorCode(statusCode) { + slog.Info("openai_image_capability_loss_skipped_by_error_code_policy", "account_id", account.ID, "status_code", statusCode) + return false + } + if !isOpenAIImageCapabilityLossError(statusCode, responseBody) { + return false + } + + resetAt := time.Now().Add(openAIImageCapabilityLossCooldown) + if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, openAIImageGenerationRateLimitKey, resetAt, openAIImageCapabilityLossReason); err != nil { + slog.Warn("openai_image_capability_loss_set_model_rate_limit_failed", "account_id", account.ID, "scope", openAIImageGenerationRateLimitKey, "error", err) + return true + } + slog.Info("openai_image_capability_lost", "account_id", account.ID, "scope", openAIImageGenerationRateLimitKey, "reset_at", resetAt, "reset_in", time.Until(resetAt).Truncate(time.Second)) + return true +} + +// isOpenAIImageCapabilityLossError reports whether upstream rejected the +// image_generation tool choice that sub2api itself put into the request body. +// Only meaningful for self-built images requests, where tools always carries a +// matching image_generation entry — upstream saying otherwise means the account +// lost the capability. +func isOpenAIImageCapabilityLossError(statusCode int, body []byte) bool { + if statusCode != http.StatusBadRequest || len(body) == 0 { + return false + } + lower := strings.ToLower(string(body)) + return strings.Contains(lower, "image_generation") && + strings.Contains(lower, "not found in 'tools' parameter") +} + func isOpenAIImageRateLimitError(statusCode int, body []byte) bool { if statusCode != http.StatusTooManyRequests || len(body) == 0 { return false diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go index 48e6a41def..c1f9f06954 100644 --- a/backend/internal/service/ratelimit_service_401_test.go +++ b/backend/internal/service/ratelimit_service_401_test.go @@ -25,6 +25,7 @@ type rateLimitAccountRepoStub struct { lastTempReason string lastErrorID int64 lastTempID int64 + tempErr error } func (r *rateLimitAccountRepoStub) SetError(ctx context.Context, id int64, errorMsg string) error { @@ -38,7 +39,7 @@ func (r *rateLimitAccountRepoStub) SetTempUnschedulable(ctx context.Context, id r.tempCalls++ r.lastTempID = id r.lastTempReason = reason - return nil + return r.tempErr } func (r *rateLimitAccountRepoStub) UpdateCredentials(ctx context.Context, id int64, credentials map[string]any) error { diff --git a/backend/internal/service/ratelimit_service_openai_image_test.go b/backend/internal/service/ratelimit_service_openai_image_test.go index 26714cf123..c62da4d80a 100644 --- a/backend/internal/service/ratelimit_service_openai_image_test.go +++ b/backend/internal/service/ratelimit_service_openai_image_test.go @@ -120,7 +120,11 @@ func TestOpenAIGatewayServiceForwardImages_ImageRateLimitReturnsFailoverAndCools require.Equal(t, openAIImageGenerationRateLimitKey, repo.modelRateLimitCalls[0].scope) } -func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *testing.T) { +// issue #6171:上游"回文字没回图"是**这一轮**的结果(模型选择了说话),不是账号能力 +// 失效。它同时被判为可重试(502)并驱动 failover,若还写 30 分钟账号级冷却,一次闲聊 +// 回复就会沿号池把每个被重试到的账号依次冷却掉。冷却仍保留给结构化上游证据,见 +// TestOpenAIGatewayServiceForwardImages_StructuredUnavailableCoolsImageCapability。 +func TestOpenAIGatewayServiceForwardImages_TextFallbackDoesNotCoolImageCapability(t *testing.T) { gin.SetMode(gin.TestMode) repo := &modelNotFoundAccountRepoStub{} body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`) @@ -154,7 +158,6 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t }, } - before := time.Now() result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") require.Nil(t, result) @@ -162,6 +165,56 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t var failoverErr *UpstreamFailoverError require.ErrorAs(t, err, &failoverErr) require.False(t, failoverErr.RetryableOnSameAccount) + // 换号行为不变:该判据仍足以放弃本账号重试这一次请求…… + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + // ……但不再写任何账号级状态,否则重试会把冷却一路刷到整个号池。 + require.Empty(t, repo.modelRateLimitCalls, + "模型回文字只说明这一轮没出图,不构成账号 30 分钟不可用的证据") +} + +// 对照不变式:上游 error 帧点名 image_generation_unavailable 时仍写冷却, +// 保证 #6171 的修复没有把这项能力保护整个废掉。 +func TestOpenAIGatewayServiceForwardImages_StructuredUnavailableCoolsImageCapability(t *testing.T) { + gin.SetMode(gin.TestMode) + repo := &modelNotFoundAccountRepoStub{} + body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`) + upstreamSSE := "data: {\"type\":\"response.failed\",\"response\":{\"id\":\"r\",\"error\":" + + "{\"type\":\"upstream_error\",\"code\":\"image_generation_unavailable\"," + + "\"message\":\"image generation tool is not available for this account\"}}}\n\n" + + 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{ + accountRepo: repo, + httpUpstream: &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }, + }, + } + parsed, err := svc.ParseOpenAIImagesRequest(c, body) + require.NoError(t, err) + account := &Account{ + ID: 206, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token-123", + }, + } + + before := time.Now() + result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") + + require.Nil(t, result) + require.Error(t, err) require.Len(t, repo.modelRateLimitCalls, 1) call := repo.modelRateLimitCalls[0] require.Equal(t, account.ID, call.accountID) @@ -169,3 +222,120 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t require.Equal(t, openAIImagesOAuthUnavailableReason, call.reason) require.WithinDuration(t, before.Add(openAIImagesOAuthUnavailableCooldown), call.resetAt, time.Second) } + +func TestOpenAIGatewayServiceForwardImages_CapabilityLossCoolsImageScope(t *testing.T) { + gin.SetMode(gin.TestMode) + repo := &modelNotFoundAccountRepoStub{} + body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`) + errorBody := `{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}` + + 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{ + rateLimitService: &RateLimitService{accountRepo: repo}, + httpUpstream: &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"X-Request-Id": []string{"req_img_capability_lost"}}, + Body: io.NopCloser(strings.NewReader(errorBody)), + }, + }, + } + parsed, err := svc.ParseOpenAIImagesRequest(c, body) + require.NoError(t, err) + account := &Account{ + ID: 205, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token-123", + }, + } + + before := time.Now() + result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") + + require.Nil(t, result) + require.Error(t, err) + require.Len(t, repo.modelRateLimitCalls, 1) + call := repo.modelRateLimitCalls[0] + require.Equal(t, account.ID, call.accountID) + require.Equal(t, openAIImageGenerationRateLimitKey, call.scope) + require.Equal(t, openAIImageCapabilityLossReason, call.reason) + require.WithinDuration(t, before.Add(openAIImageCapabilityLossCooldown), call.resetAt, time.Second) +} + +func TestOpenAIGatewayServiceHandleUpstreamError_PassthroughCapabilityLossDoesNotCool(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}} + account := &Account{ID: 206, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + body := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`) + + disabled := svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, http.Header{}, body, "gpt-5.5") + + require.False(t, disabled) + require.Empty(t, repo.modelRateLimitCalls) + _, wholeAccountBlocked := svc.openaiAccountRuntimeBlockUntil.Load(account.ID) + require.False(t, wholeAccountBlocked) +} + +func TestRateLimitServiceHandleOpenAIImageCapabilityLoss_IgnoresGenericBadRequest(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := &Account{ID: 207, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + body := []byte(`{"error":{"message":"Invalid type for input[0].arguments"}}`) + + handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body) + + require.False(t, handled) + require.Empty(t, repo.modelRateLimitCalls) +} + +func TestRateLimitServiceHandleOpenAIImageCapabilityLoss_RespectsPlatformAndErrorCodePolicy(t *testing.T) { + body := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`) + + t.Run("non_openai_platform", func(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := &Account{ID: 208, Platform: PlatformAnthropic, Type: AccountTypeOAuth} + + handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body) + + require.False(t, handled) + require.Empty(t, repo.modelRateLimitCalls) + }) + + t.Run("custom_error_code_policy_excludes_400", func(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := &Account{ + ID: 209, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "custom_error_codes_enabled": true, + "custom_error_codes": []any{float64(http.StatusTooManyRequests)}, + }, + } + + require.False(t, account.ShouldHandleErrorCode(http.StatusBadRequest)) + handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body) + + require.False(t, handled) + require.Empty(t, repo.modelRateLimitCalls) + }) +} + +func TestIsOpenAIImageCapabilityLossError(t *testing.T) { + capabilityLossBody := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`) + genericBadRequestBody := []byte(`{"error":{"message":"Invalid type for input[0].arguments"}}`) + + require.True(t, isOpenAIImageCapabilityLossError(http.StatusBadRequest, capabilityLossBody)) + require.False(t, isOpenAIImageCapabilityLossError(http.StatusBadRequest, genericBadRequestBody)) + require.False(t, isOpenAIImageCapabilityLossError(http.StatusTooManyRequests, capabilityLossBody)) +} diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index 790b2243d2..d3d88c4269 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -3,17 +3,128 @@ package service import ( "context" "encoding/json" + "errors" "fmt" "io" + "log/slog" "net/http" + "net/url" "sort" "strings" + "time" "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" "github.com/Wei-Shaw/sub2api/internal/pkg/claude" "github.com/Wei-Shaw/sub2api/internal/pkg/geminicli" ) +const ( + upstreamModelsBodyLimit int64 = 8 << 20 + modelsDevRegistryURL = "https://models.dev/api.json" + modelsDevRegistryTTL = 6 * time.Hour + UpstreamModelMetadataExtraKey = "upstream_model_metadata" + UpstreamModelMetadataIncompleteCode = "upstream_model_metadata_incomplete" +) + +type UpstreamModelMetadata struct { + ID string `json:"id"` + DisplayName string `json:"display_name,omitempty"` + Description string `json:"description,omitempty"` + Reasoning *bool `json:"reasoning,omitempty"` + DefaultReasoningLevel string `json:"default_reasoning_level,omitempty"` + SupportedReasoningLevels []string `json:"supported_reasoning_levels,omitempty"` + InputModalities []string `json:"input_modalities,omitempty"` + ContextWindow int64 `json:"context_window,omitempty"` + MaxOutputTokens int64 `json:"max_output_tokens,omitempty"` +} + +type UpstreamModelMetadataSnapshot struct { + Source string `json:"source"` + SyncedAt string `json:"synced_at"` + Models map[string]UpstreamModelMetadata `json:"models"` +} + +type UpstreamModelCatalog struct { + Models []string `json:"models"` + Metadata map[string]UpstreamModelMetadata `json:"metadata,omitempty"` + Warnings []UpstreamModelSyncWarning `json:"warnings,omitempty"` +} + +type UpstreamModelSyncWarning struct { + Code string `json:"code"` + Message string `json:"message"` +} + +type modelsDevProvider struct { + ID string `json:"id"` + Name string `json:"name"` + API string `json:"api"` + Models map[string]modelsDevModel `json:"models"` +} + +type modelsDevModel struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description"` + Reasoning *bool `json:"reasoning"` + ReasoningOptions []modelsDevReasoningOption `json:"reasoning_options"` + Modalities modelsDevModalities `json:"modalities"` + Limit modelsDevLimit `json:"limit"` +} + +type modelsDevReasoningOption struct { + Type string `json:"type"` + Values []any `json:"values"` +} + +type modelsDevModalities struct { + Input []string `json:"input"` + Output []string `json:"output"` +} + +type modelsDevLimit struct { + Context int64 `json:"context"` + Output int64 `json:"output"` +} + +func (a *Account) SetUpstreamModelMetadataSnapshot(snapshot UpstreamModelMetadataSnapshot) { + if a == nil { + return + } + if a.Extra == nil { + a.Extra = make(map[string]any) + } + a.Extra[UpstreamModelMetadataExtraKey] = snapshot +} + +func (a *Account) GetUpstreamModelMetadataSnapshot() *UpstreamModelMetadataSnapshot { + if a == nil || a.Extra == nil { + return nil + } + raw, ok := a.Extra[UpstreamModelMetadataExtraKey] + if !ok || raw == nil { + return nil + } + body, err := json.Marshal(raw) + if err != nil { + return nil + } + var snapshot UpstreamModelMetadataSnapshot + if err := json.Unmarshal(body, &snapshot); err != nil || len(snapshot.Models) == 0 { + return nil + } + return &snapshot +} + +func (a *Account) GetUpstreamModelMetadata(modelID string) (UpstreamModelMetadata, bool) { + snapshot := a.GetUpstreamModelMetadataSnapshot() + if snapshot == nil { + return UpstreamModelMetadata{}, false + } + metadata, ok := snapshot.Models[strings.TrimSpace(modelID)] + return metadata, ok +} + // UpstreamModelSyncErrorKind classifies model sync failures for safe HTTP mapping. type UpstreamModelSyncErrorKind string @@ -24,13 +135,16 @@ const ( UpstreamModelSyncErrorUnsupported UpstreamModelSyncErrorKind = "unsupported" // UpstreamModelSyncErrorUpstream means the configured upstream failed or returned an unusable response. UpstreamModelSyncErrorUpstream UpstreamModelSyncErrorKind = "upstream" + // UpstreamModelSyncErrorInternal means local persistence or service state failed after a valid upstream response. + UpstreamModelSyncErrorInternal UpstreamModelSyncErrorKind = "internal" ) // UpstreamModelSyncError keeps internal failure details wrapped while exposing a safe client message. type UpstreamModelSyncError struct { - Kind UpstreamModelSyncErrorKind - Message string - Err error + Kind UpstreamModelSyncErrorKind + Message string + StatusCode int + Err error } func (e *UpstreamModelSyncError) Error() string { @@ -70,49 +184,432 @@ func newUpstreamModelSyncUpstreamError(message string, err error) error { return &UpstreamModelSyncError{Kind: UpstreamModelSyncErrorUpstream, Message: message, Err: err} } -// FetchUpstreamSupportedModels fetches the live model list from the account's upstream API format. +func newUpstreamModelSyncInternalError(message string, err error) error { + return &UpstreamModelSyncError{Kind: UpstreamModelSyncErrorInternal, Message: message, Err: err} +} + +// FetchUpstreamSupportedModels fetches only live model IDs. The admin sync path +// uses SyncUpstreamModelCatalog so capability metadata can also be persisted. func (s *AccountTestService) FetchUpstreamSupportedModels(ctx context.Context, account *Account) ([]string, error) { + models, _, err := s.fetchUpstreamModelList(ctx, account) + return models, err +} + +// SyncUpstreamModelCatalog fetches the account's live model list, enriches +// missing capability fields from the provider registry used by the upstream, +// and persists a normalized account snapshot when metadata is available. +func (s *AccountTestService) SyncUpstreamModelCatalog(ctx context.Context, account *Account) (*UpstreamModelCatalog, error) { + models, body, err := s.fetchUpstreamModelList(ctx, account) + if err != nil { + configuredModels := configuredUpstreamModelsForCapabilitySync(account) + if !upstreamModelListEndpointUnsupported(err) || len(configuredModels) == 0 { + return nil, err + } + models = configuredModels + body = nil + slog.Info("upstream model list endpoint unavailable; using configured models for capability sync", + "account_id", upstreamModelSyncAccountID(account), + "platform", upstreamModelSyncPlatform(account), + "status_code", upstreamModelSyncStatusCode(err), + "model_count", len(models), + ) + } + catalog := &UpstreamModelCatalog{Models: models, Metadata: make(map[string]UpstreamModelMetadata)} + if len(body) > 0 { + _, directMetadata, parseErr := extractUpstreamModelCatalog(body, account != nil && account.IsGrok()) + if parseErr == nil { + catalog.Metadata = directMetadata + } + } + + source := "upstream" + metadataIncomplete := upstreamCatalogNeedsRegistry(models, catalog.Metadata) + if metadataIncomplete { + if registryMetadata, registryErr := s.fetchModelsDevMetadata(ctx, account, models); registryErr == nil { + for modelID, fallback := range registryMetadata { + current := catalog.Metadata[modelID] + merged, changed := mergeUpstreamModelMetadata(current, fallback) + catalog.Metadata[modelID] = merged + if changed { + source = "models.dev" + } + } + } else { + slog.Warn("upstream model capability metadata enrichment failed", + "account_id", upstreamModelSyncAccountID(account), + "platform", upstreamModelSyncPlatform(account), + "error", registryErr, + ) + } + } + + if upstreamCatalogNeedsRegistry(models, catalog.Metadata) { + catalog.Warnings = append(catalog.Warnings, UpstreamModelSyncWarning{ + Code: UpstreamModelMetadataIncompleteCode, + Message: "Model IDs were synced, but capability metadata is incomplete.", + }) + return catalog, nil + } + if len(catalog.Metadata) == 0 || account == nil || account.ID <= 0 || s.accountRepo == nil { + return catalog, nil + } + snapshot := UpstreamModelMetadataSnapshot{ + Source: source, + SyncedAt: time.Now().UTC().Format(time.RFC3339), + Models: catalog.Metadata, + } + if err := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{UpstreamModelMetadataExtraKey: snapshot}); err != nil { + return nil, newUpstreamModelSyncInternalError("Failed to save upstream model metadata", err) + } + account.SetUpstreamModelMetadataSnapshot(snapshot) + return catalog, nil +} + +func upstreamModelSyncStatusCode(err error) int { + var syncErr *UpstreamModelSyncError + if errors.As(err, &syncErr) { + return syncErr.StatusCode + } + return 0 +} + +func upstreamModelListEndpointUnsupported(err error) bool { + statusCode := upstreamModelSyncStatusCode(err) + return statusCode == http.StatusNotFound || statusCode == http.StatusMethodNotAllowed +} + +func configuredUpstreamModelsForCapabilitySync(account *Account) []string { + if account == nil { + return nil + } + models := make([]string, 0) + for _, mappedModel := range account.GetModelMapping() { + mappedModel = strings.TrimSpace(mappedModel) + if mappedModel == "" || strings.Contains(mappedModel, "*") { + continue + } + models = append(models, mappedModel) + } + return dedupeAndSortModelIDs(models) +} + +func upstreamModelSyncAccountID(account *Account) int64 { + if account == nil { + return 0 + } + return account.ID +} + +func upstreamModelSyncPlatform(account *Account) string { + if account == nil { + return "" + } + return account.Platform +} + +func upstreamCatalogNeedsRegistry(models []string, metadata map[string]UpstreamModelMetadata) bool { + for _, modelID := range models { + modelID = strings.TrimSpace(modelID) + model, ok := metadata[modelID] + if !ok || !upstreamModelMetadataIsUseful(model) { + return true + } + if model.Reasoning == nil || len(model.InputModalities) == 0 || model.ContextWindow <= 0 { + return true + } + if *model.Reasoning && len(model.SupportedReasoningLevels) == 0 { + return true + } + } + return false +} + +func upstreamModelMetadataIsUseful(metadata UpstreamModelMetadata) bool { + return strings.TrimSpace(metadata.DisplayName) != "" || + strings.TrimSpace(metadata.Description) != "" || + metadata.Reasoning != nil || + len(metadata.SupportedReasoningLevels) > 0 || + len(metadata.InputModalities) > 0 || + metadata.ContextWindow > 0 || + metadata.MaxOutputTokens > 0 +} + +func mergeUpstreamModelMetadata(primary, fallback UpstreamModelMetadata) (UpstreamModelMetadata, bool) { + merged := primary + changed := false + if strings.TrimSpace(merged.ID) == "" && strings.TrimSpace(fallback.ID) != "" { + merged.ID = strings.TrimSpace(fallback.ID) + changed = true + } + if strings.TrimSpace(merged.DisplayName) == "" && strings.TrimSpace(fallback.DisplayName) != "" { + merged.DisplayName = strings.TrimSpace(fallback.DisplayName) + changed = true + } + if strings.TrimSpace(merged.Description) == "" && strings.TrimSpace(fallback.Description) != "" { + merged.Description = strings.TrimSpace(fallback.Description) + changed = true + } + if merged.Reasoning == nil && fallback.Reasoning != nil { + reasoning := *fallback.Reasoning + merged.Reasoning = &reasoning + changed = true + } + if strings.TrimSpace(merged.DefaultReasoningLevel) == "" && strings.TrimSpace(fallback.DefaultReasoningLevel) != "" { + merged.DefaultReasoningLevel = strings.TrimSpace(fallback.DefaultReasoningLevel) + changed = true + } + if len(merged.SupportedReasoningLevels) == 0 && len(fallback.SupportedReasoningLevels) > 0 { + merged.SupportedReasoningLevels = append([]string(nil), fallback.SupportedReasoningLevels...) + changed = true + } + if len(merged.InputModalities) == 0 && len(fallback.InputModalities) > 0 { + merged.InputModalities = append([]string(nil), fallback.InputModalities...) + changed = true + } + if merged.ContextWindow <= 0 && fallback.ContextWindow > 0 { + merged.ContextWindow = fallback.ContextWindow + changed = true + } + if merged.MaxOutputTokens <= 0 && fallback.MaxOutputTokens > 0 { + merged.MaxOutputTokens = fallback.MaxOutputTokens + changed = true + } + return merged, changed +} + +func (s *AccountTestService) fetchModelsDevMetadata( + ctx context.Context, + account *Account, + modelIDs []string, +) (map[string]UpstreamModelMetadata, error) { + if s == nil || s.httpUpstream == nil || account == nil { + return nil, fmt.Errorf("model metadata registry is not configured") + } + registry, err := s.fetchModelsDevRegistry(ctx, account) + if err != nil { + return nil, err + } + provider, ok := matchModelsDevProvider(registry, upstreamModelRegistryBaseURL(account)) + if !ok { + return nil, fmt.Errorf("no model metadata provider matches account base URL") + } + + metadata := make(map[string]UpstreamModelMetadata) + for _, modelID := range modelIDs { + modelID = strings.TrimSpace(modelID) + model, found := provider.Models[modelID] + if !found { + for candidateID, candidate := range provider.Models { + if strings.EqualFold(strings.TrimSpace(candidateID), modelID) || strings.EqualFold(strings.TrimSpace(candidate.ID), modelID) { + model = candidate + found = true + break + } + } + } + if !found { + continue + } + entry := upstreamMetadataFromModelsDevModel(modelID, model) + if upstreamModelMetadataIsUseful(entry) { + metadata[modelID] = entry + } + } + return metadata, nil +} + +func (s *AccountTestService) fetchModelsDevRegistry(ctx context.Context, account *Account) (map[string]modelsDevProvider, error) { + now := time.Now() + s.modelMetadataRegistryMu.Lock() + if len(s.modelMetadataRegistry) > 0 && now.Sub(s.modelMetadataRegistryAt) < modelsDevRegistryTTL { + cached := s.modelMetadataRegistry + s.modelMetadataRegistryMu.Unlock() + return cached, nil + } + s.modelMetadataRegistryMu.Unlock() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, modelsDevRegistryURL, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/json") + resp, err := s.doUpstreamModelsRequest(req, upstreamModelsProxyURL(account), account) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return nil, fmt.Errorf("model metadata registry returned HTTP %d", resp.StatusCode) + } + body, err := io.ReadAll(io.LimitReader(resp.Body, upstreamModelsBodyLimit+1)) + if err != nil { + return nil, err + } + if int64(len(body)) > upstreamModelsBodyLimit { + return nil, fmt.Errorf("model metadata registry response exceeds %d bytes", upstreamModelsBodyLimit) + } + var registry map[string]modelsDevProvider + if err := json.Unmarshal(body, ®istry); err != nil { + return nil, fmt.Errorf("parse model metadata registry: %w", err) + } + if len(registry) == 0 { + return nil, fmt.Errorf("model metadata registry is empty") + } + + s.modelMetadataRegistryMu.Lock() + s.modelMetadataRegistry = registry + s.modelMetadataRegistryAt = now + s.modelMetadataRegistryMu.Unlock() + return registry, nil +} + +func upstreamMetadataFromModelsDevModel(modelID string, model modelsDevModel) UpstreamModelMetadata { + levels := reasoningLevelsFromModelsDevOptions(model.ReasoningOptions) + reasoning := model.Reasoning + if reasoning == nil && len(levels) > 0 { + inferred := true + reasoning = &inferred + } + metadata := UpstreamModelMetadata{ + ID: strings.TrimSpace(modelID), + DisplayName: strings.TrimSpace(model.Name), + Description: strings.TrimSpace(model.Description), + Reasoning: reasoning, + SupportedReasoningLevels: levels, + InputModalities: normalizeCodexInputModalities(model.Modalities.Input), + ContextWindow: model.Limit.Context, + MaxOutputTokens: model.Limit.Output, + } + if len(levels) > 0 { + metadata.DefaultReasoningLevel = levels[0] + } + if strings.TrimSpace(model.ID) != "" { + metadata.ID = strings.TrimSpace(model.ID) + } + return metadata +} + +func reasoningLevelsFromModelsDevOptions(options []modelsDevReasoningOption) []string { + levels := make([]string, 0) + for _, option := range options { + if !strings.EqualFold(strings.TrimSpace(option.Type), "effort") { + continue + } + for _, value := range option.Values { + if value == nil { + levels = append(levels, "none") + continue + } + if effort, ok := value.(string); ok { + levels = append(levels, effort) + } + } + } + return normalizeReasoningLevels(levels) +} + +func upstreamModelRegistryBaseURL(account *Account) string { + if account == nil { + return "" + } + switch { + case account.IsOpenAI() || account.IsCNProvider(): + return account.GetOpenAIFormatBaseURL() + case account.IsGrok(): + return account.GetGrokBaseURL() + case account.IsGemini(): + return account.GetGeminiBaseURL(geminicli.AIStudioBaseURL) + case account.IsAnthropic(): + return account.GetBaseURL() + case account.Platform == PlatformAntigravity: + return account.GetGeminiBaseURL(geminicli.AIStudioBaseURL) + default: + return strings.TrimSpace(account.GetCredential("base_url")) + } +} + +func matchModelsDevProvider(registry map[string]modelsDevProvider, accountBaseURL string) (modelsDevProvider, bool) { + accountBaseURL = normalizeModelRegistryBaseURL(accountBaseURL) + if accountBaseURL == "" { + return modelsDevProvider{}, false + } + var best modelsDevProvider + bestScore := -1 + for _, provider := range registry { + providerBaseURL := normalizeModelRegistryBaseURL(provider.API) + if providerBaseURL == "" { + continue + } + if accountBaseURL != providerBaseURL && + !strings.HasPrefix(accountBaseURL, providerBaseURL+"/") && + !strings.HasPrefix(providerBaseURL, accountBaseURL+"/") { + continue + } + if len(providerBaseURL) > bestScore { + best = provider + bestScore = len(providerBaseURL) + } + } + return best, bestScore >= 0 +} + +func normalizeModelRegistryBaseURL(raw string) string { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || parsed.Scheme == "" || parsed.Host == "" { + return "" + } + path := strings.TrimRight(parsed.Path, "/") + if strings.HasSuffix(strings.ToLower(path), "/models") { + path = strings.TrimRight(path[:len(path)-len("/models")], "/") + } + return strings.ToLower(parsed.Scheme) + "://" + strings.ToLower(parsed.Host) + path +} + +func (s *AccountTestService) fetchUpstreamModelList(ctx context.Context, account *Account) ([]string, []byte, error) { if s == nil { - return nil, newUpstreamModelSyncConfigError("Account test service is not configured", nil) + return nil, nil, newUpstreamModelSyncConfigError("Account test service is not configured", nil) } if account == nil { - return nil, newUpstreamModelSyncConfigError("Account is required", nil) + return nil, nil, newUpstreamModelSyncConfigError("Account is required", nil) } if account.Platform == PlatformAntigravity && account.Type != AccountTypeAPIKey { - return s.fetchAntigravityOAuthUpstreamModels(ctx, account) + models, err := s.fetchAntigravityOAuthUpstreamModels(ctx, account) + return models, nil, err } if s.httpUpstream == nil { - return nil, newUpstreamModelSyncConfigError("Upstream HTTP client is not configured", nil) + return nil, nil, newUpstreamModelSyncConfigError("Upstream HTTP client is not configured", nil) } req, err := s.buildUpstreamModelsRequest(ctx, account) if err != nil { - return nil, err + return nil, nil, err } proxyURL := upstreamModelsProxyURL(account) resp, err := s.doUpstreamModelsRequest(req, proxyURL, account) if err != nil { - return nil, newUpstreamModelSyncUpstreamError("Failed to request upstream model list", err) + return nil, nil, newUpstreamModelSyncUpstreamError("Failed to request upstream model list", err) } defer func() { _ = resp.Body.Close() }() bodyLimit := resolveModelsListReadLimit(s.cfg) body, err := io.ReadAll(io.LimitReader(resp.Body, bodyLimit+1)) if err != nil { - return nil, newUpstreamModelSyncUpstreamError("Failed to read upstream model list", err) + return nil, nil, newUpstreamModelSyncUpstreamError("Failed to read upstream model list", err) } if int64(len(body)) > bodyLimit { - return nil, newUpstreamModelSyncUpstreamError("Upstream model list response is too large", fmt.Errorf("response exceeds %d bytes", bodyLimit)) + return nil, nil, newUpstreamModelSyncUpstreamError("Upstream model list response is too large", fmt.Errorf("response exceeds %d bytes", bodyLimit)) } if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { - return nil, newUpstreamModelSyncUpstreamError( - fmt.Sprintf("Upstream model list request failed with HTTP %d", resp.StatusCode), - fmt.Errorf("upstream model list returned HTTP %d", resp.StatusCode), - ) + return nil, nil, &UpstreamModelSyncError{ + Kind: UpstreamModelSyncErrorUpstream, + Message: fmt.Sprintf("Upstream model list request failed with HTTP %d", resp.StatusCode), + StatusCode: resp.StatusCode, + Err: fmt.Errorf("upstream model list returned HTTP %d", resp.StatusCode), + } } extractModels := extractUpstreamModelIDs @@ -121,13 +618,13 @@ func (s *AccountTestService) FetchUpstreamSupportedModels(ctx context.Context, a } models, err := extractModels(body) if err != nil { - return nil, newUpstreamModelSyncUpstreamError("Upstream model list response was not valid JSON", err) + return nil, nil, newUpstreamModelSyncUpstreamError("Upstream model list response was not valid JSON", err) } if len(models) == 0 { - return nil, newUpstreamModelSyncUpstreamError("Upstream returned no supported models", nil) + return nil, nil, newUpstreamModelSyncUpstreamError("Upstream returned no supported models", nil) } - return models, nil + return models, body, nil } func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) { @@ -577,12 +1074,29 @@ type upstreamModelEntry struct { type upstreamModelEntryMetadata struct { ID string `json:"id"` + Slug string `json:"slug"` Model string `json:"model"` ModelID string `json:"modelId"` ModelIDSnake string `json:"model_id"` Name string `json:"name"` } +type upstreamModelCapabilityEntry struct { + upstreamModelEntry + DisplayName string `json:"display_name"` + Description string `json:"description"` + Reasoning *bool `json:"reasoning"` + DefaultReasoningLevel string `json:"default_reasoning_level"` + SupportedReasoningLevels []json.RawMessage `json:"supported_reasoning_levels"` + ReasoningOptions []modelsDevReasoningOption `json:"reasoning_options"` + InputModalities []string `json:"input_modalities"` + Modalities modelsDevModalities `json:"modalities"` + ContextWindow int64 `json:"context_window"` + MaxContextWindow int64 `json:"max_context_window"` + MaxOutputTokens int64 `json:"max_output_tokens"` + Limit modelsDevLimit `json:"limit"` +} + func extractUpstreamModelIDs(body []byte) ([]string, error) { return extractUpstreamModelIDsWithSelector(body, upstreamModelEntryID) } @@ -591,6 +1105,166 @@ func extractGrokUpstreamModelIDs(body []byte) ([]string, error) { return extractUpstreamModelIDsWithSelector(body, grokUpstreamModelEntryID) } +func extractUpstreamModelCatalog(body []byte, grok bool) ([]string, map[string]UpstreamModelMetadata, error) { + entries, err := extractUpstreamModelRawEntries(body) + if err != nil { + return nil, nil, err + } + selectID := upstreamModelEntryID + if grok { + selectID = grokUpstreamModelEntryID + } + + models := make([]string, 0, len(entries)) + metadata := make(map[string]UpstreamModelMetadata) + for _, raw := range entries { + var capability upstreamModelCapabilityEntry + if err := json.Unmarshal(raw, &capability); err != nil { + continue + } + modelID := strings.TrimSpace(selectID(capability.upstreamModelEntry)) + if modelID == "" { + continue + } + models = append(models, modelID) + entry := upstreamMetadataFromCapabilityEntry(modelID, capability) + if upstreamModelMetadataIsUseful(entry) { + metadata[modelID] = entry + } + } + return dedupeAndSortModelIDs(models), metadata, nil +} + +func extractUpstreamModelRawEntries(body []byte) ([]json.RawMessage, error) { + var response struct { + Data []json.RawMessage `json:"data"` + Models []json.RawMessage `json:"models"` + } + if err := json.Unmarshal(body, &response); err == nil && (response.Data != nil || response.Models != nil) { + entries := make([]json.RawMessage, 0, len(response.Data)+len(response.Models)) + entries = append(entries, response.Data...) + entries = append(entries, response.Models...) + return entries, nil + } + var entries []json.RawMessage + if err := json.Unmarshal(body, &entries); err != nil { + return nil, fmt.Errorf("parse upstream model catalog: %w", err) + } + return entries, nil +} + +func upstreamMetadataFromCapabilityEntry(modelID string, entry upstreamModelCapabilityEntry) UpstreamModelMetadata { + levels := reasoningLevelsFromRawEntries(entry.SupportedReasoningLevels) + if len(levels) == 0 { + levels = reasoningLevelsFromModelsDevOptions(entry.ReasoningOptions) + } + reasoning := entry.Reasoning + if reasoning == nil && len(levels) > 0 { + inferred := len(levels) != 1 || levels[0] != "none" + reasoning = &inferred + } + modalities := entry.InputModalities + if len(modalities) == 0 { + modalities = entry.Modalities.Input + } + contextWindow := entry.ContextWindow + if contextWindow <= 0 { + contextWindow = entry.MaxContextWindow + } + if contextWindow <= 0 { + contextWindow = entry.Limit.Context + } + maxOutputTokens := entry.MaxOutputTokens + if maxOutputTokens <= 0 { + maxOutputTokens = entry.Limit.Output + } + defaultReasoningLevel := normalizeReasoningLevel(entry.DefaultReasoningLevel) + if defaultReasoningLevel == "" && len(levels) > 0 { + defaultReasoningLevel = levels[0] + } + displayName := strings.TrimSpace(entry.DisplayName) + if displayName == "" && strings.TrimSpace(entry.Name) != "" && strings.TrimSpace(entry.Name) != modelID { + displayName = strings.TrimSpace(entry.Name) + } + return UpstreamModelMetadata{ + ID: modelID, + DisplayName: displayName, + Description: strings.TrimSpace(entry.Description), + Reasoning: reasoning, + DefaultReasoningLevel: defaultReasoningLevel, + SupportedReasoningLevels: levels, + InputModalities: normalizeCodexInputModalities(modalities), + ContextWindow: contextWindow, + MaxOutputTokens: maxOutputTokens, + } +} + +func reasoningLevelsFromRawEntries(entries []json.RawMessage) []string { + levels := make([]string, 0, len(entries)) + for _, raw := range entries { + var effort string + if err := json.Unmarshal(raw, &effort); err == nil { + levels = append(levels, effort) + continue + } + var level struct { + Effort string `json:"effort"` + } + if err := json.Unmarshal(raw, &level); err == nil { + levels = append(levels, level.Effort) + } + } + return normalizeReasoningLevels(levels) +} + +func normalizeReasoningLevels(levels []string) []string { + seen := make(map[string]struct{}, len(levels)) + normalized := make([]string, 0, len(levels)) + for _, level := range levels { + level = normalizeReasoningLevel(level) + if level == "" { + continue + } + if _, exists := seen[level]; exists { + continue + } + seen[level] = struct{}{} + normalized = append(normalized, level) + } + return normalized +} + +func normalizeReasoningLevel(level string) string { + level = strings.ToLower(strings.TrimSpace(level)) + switch level { + case "off", "disabled": + return "none" + case "extra-high", "extra_high": + return "xhigh" + case "none", "minimal", "low", "medium", "high", "xhigh", "max", "ultra": + return level + default: + return "" + } +} + +func normalizeCodexInputModalities(modalities []string) []string { + seen := make(map[string]struct{}, len(modalities)) + normalized := make([]string, 0, len(modalities)) + for _, modality := range modalities { + modality = strings.ToLower(strings.TrimSpace(modality)) + if modality != "text" && modality != "image" { + continue + } + if _, exists := seen[modality]; exists { + continue + } + seen[modality] = struct{}{} + normalized = append(normalized, modality) + } + return normalized +} + func extractUpstreamModelIDsWithSelector(body []byte, selectID func(upstreamModelEntry) string) ([]string, error) { var response struct { Data []upstreamModelEntry `json:"data"` @@ -646,6 +1320,7 @@ func grokUpstreamModelEntryID(entry upstreamModelEntry) string { entry.ModelID, entry.ModelIDSnake, entry.ID, + entry.Slug, } if len(entry.Meta) > 0 { var meta upstreamModelEntryMetadata @@ -655,6 +1330,7 @@ func grokUpstreamModelEntryID(entry upstreamModelEntry) string { meta.ModelID, meta.ModelIDSnake, meta.ID, + meta.Slug, meta.Name, ) } diff --git a/backend/internal/service/upstream_models_test.go b/backend/internal/service/upstream_models_test.go index 114c647d45..b358406b8b 100644 --- a/backend/internal/service/upstream_models_test.go +++ b/backend/internal/service/upstream_models_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "encoding/json" "errors" "io" "net/http" @@ -13,6 +14,28 @@ import ( "github.com/stretchr/testify/require" ) +type upstreamModelMetadataRepoStub struct { + AccountRepository + accountID int64 + updates map[string]any + err error +} + +func headerValuesEqualFold(header http.Header, name string) []string { + for key, values := range header { + if strings.EqualFold(key, name) { + return values + } + } + return nil +} + +func (r *upstreamModelMetadataRepoStub) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + r.accountID = id + r.updates = updates + return r.err +} + func upstreamModelSyncTestConfig() *config.Config { return &config.Config{ Security: config.SecurityConfig{ @@ -394,6 +417,384 @@ func TestFetchUpstreamSupportedModelsParsesOpenAIResponse(t *testing.T) { require.Equal(t, "Bearer openai-key", upstream.lastReq.Header.Get("Authorization")) } +// Scenario: ID-only 模型列表从 Models.dev 补齐能力。 +func TestSyncUpstreamModelCatalogEnrichesOpenCodeIDOnlyListAndPersistsSnapshot(t *testing.T) { + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"object":"list","data":[{"id":"x-preview-f-free","object":"model"}]}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "opencode": { + "id": "opencode", + "name": "OpenCode Zen", + "api": "https://opencode.ai/zen/v1", + "models": { + "x-preview-f-free": { + "id": "x-preview-f-free", + "name": "Ox Alpha Free (Unlimited)", + "description": "Stealth reasoning model for coding, agentic tasks, and tool use", + "reasoning": true, + "reasoning_options": [{"type":"effort","values":["low","high","max"]}], + "modalities": {"input":["text","image","video"],"output":["text"]}, + "limit": {"context":1000000,"output":131072} + } + } + } + }`)), + }, + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{ + accountRepo: repo, + httpUpstream: upstream, + cfg: upstreamModelSyncTestConfig(), + } + account := &Account{ + ID: 91, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "opencode-key", + "base_url": "https://opencode.ai/zen/v1", + "header_override_enabled": true, + "header_overrides": map[string]any{ + "X-Custom-Account-Header": "account-secret", + }, + }, + } + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), account) + require.NoError(t, err) + require.Equal(t, []string{"x-preview-f-free"}, catalog.Models) + require.Len(t, upstream.requests, 2) + require.Equal(t, "https://opencode.ai/zen/v1/models", upstream.requests[0].URL.String()) + require.Equal(t, []string{"account-secret"}, headerValuesEqualFold(upstream.requests[0].Header, "X-Custom-Account-Header")) + require.Equal(t, modelsDevRegistryURL, upstream.requests[1].URL.String()) + require.Empty(t, upstream.requests[1].Header.Get("Authorization")) + require.Empty(t, upstream.requests[1].Header.Get("x-api-key")) + require.Empty(t, headerValuesEqualFold(upstream.requests[1].Header, "X-Custom-Account-Header")) + + metadata := catalog.Metadata["x-preview-f-free"] + require.Equal(t, "Ox Alpha Free (Unlimited)", metadata.DisplayName) + require.NotNil(t, metadata.Reasoning) + require.True(t, *metadata.Reasoning) + require.Equal(t, []string{"low", "high", "max"}, metadata.SupportedReasoningLevels) + require.Equal(t, []string{"text", "image"}, metadata.InputModalities) + require.Equal(t, int64(1_000_000), metadata.ContextWindow) + require.Equal(t, int64(131_072), metadata.MaxOutputTokens) + require.Equal(t, int64(91), repo.accountID) + + rawSnapshot, ok := repo.updates[UpstreamModelMetadataExtraKey] + require.True(t, ok) + encoded, err := json.Marshal(rawSnapshot) + require.NoError(t, err) + var snapshot UpstreamModelMetadataSnapshot + require.NoError(t, json.Unmarshal(encoded, &snapshot)) + require.Equal(t, "models.dev", snapshot.Source) + require.Equal(t, metadata, snapshot.Models["x-preview-f-free"]) +} + +// Scenario: 不提供 /models 的兼容上游使用管理员已配置模型继续同步能力。 +func TestSyncUpstreamModelCatalogUsesConfiguredModelsWhenListEndpointUnsupported(t *testing.T) { + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusNotFound, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":"not found"}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "configured-provider": { + "id": "configured-provider", + "name": "Configured Provider", + "api": "https://provider.example/v1", + "models": { + "glm-5.3": { + "id": "glm-5.3", + "name": "GLM-5.3", + "reasoning": true, + "reasoning_options": [{"type":"effort","values":["low","medium","high"]}], + "modalities": {"input":["text"],"output":["text"]}, + "limit": {"context":1000000,"output":131072} + } + } + } + }`)), + }, + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + account := &Account{ + ID: 97, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "key", + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{ + "public-glm": "glm-5.3", + "duplicate": "glm-5.3", + "wildcard": "glm-*", + "empty": "", + }, + }, + } + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), account) + require.NoError(t, err) + require.Equal(t, []string{"glm-5.3"}, catalog.Models) + require.Empty(t, catalog.Warnings) + require.Len(t, upstream.requests, 2) + require.Equal(t, "https://provider.example/v1/models", upstream.requests[0].URL.String()) + require.Equal(t, modelsDevRegistryURL, upstream.requests[1].URL.String()) + metadata := catalog.Metadata["glm-5.3"] + require.Equal(t, []string{"low", "medium", "high"}, metadata.SupportedReasoningLevels) + require.Equal(t, []string{"text"}, metadata.InputModalities) + require.Equal(t, int64(1_000_000), metadata.ContextWindow) + require.NotNil(t, repo.updates) +} + +func TestSyncUpstreamModelCatalogDoesNotUseConfiguredModelsForRealUpstreamFailures(t *testing.T) { + tests := []struct { + name string + statusCode int + }{ + {name: "unauthorized", statusCode: http.StatusUnauthorized}, + {name: "rate limited", statusCode: http.StatusTooManyRequests}, + {name: "server error", statusCode: http.StatusBadGateway}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: tt.statusCode, + Body: io.NopCloser(strings.NewReader(`{"error":"failed"}`)), + }} + svc := &AccountTestService{httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + _, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 98, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "key", + "base_url": "https://provider.example/v1", + "model_mapping": map[string]any{"public-glm": "glm-5.3"}, + }, + }) + require.Error(t, err) + require.Len(t, upstream.requests, 1) + require.Equal(t, tt.statusCode, upstreamModelSyncStatusCode(err)) + }) + } +} + +func TestSyncUpstreamModelCatalogRequiresConfiguredModelsForUnsupportedListEndpoint(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusMethodNotAllowed, + Body: io.NopCloser(strings.NewReader(`{"error":"method not allowed"}`)), + }} + svc := &AccountTestService{httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + _, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 99, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + }) + require.Error(t, err) + require.Equal(t, http.StatusMethodNotAllowed, upstreamModelSyncStatusCode(err)) + require.Len(t, upstream.requests, 1) +} + +// Scenario: 完整上游模型清单优先保存能力。 +func TestSyncUpstreamModelCatalogPrefersDirectUpstreamMetadata(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"models":[{ + "slug":"custom-thinking-model", + "display_name":"Upstream Display", + "description":"Upstream description", + "default_reasoning_level":"high", + "supported_reasoning_levels":[{"effort":"low"},{"effort":"high"},{"effort":"ultra"}], + "input_modalities":["text","image"], + "context_window":256000 + }]}`)), + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 92, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + }) + require.NoError(t, err) + require.Len(t, upstream.requests, 1, "complete upstream metadata must not be replaced by a registry fetch") + metadata := catalog.Metadata["custom-thinking-model"] + require.Equal(t, "Upstream Display", metadata.DisplayName) + require.Equal(t, "high", metadata.DefaultReasoningLevel) + require.Equal(t, []string{"low", "high", "ultra"}, metadata.SupportedReasoningLevels) + require.Equal(t, []string{"text", "image"}, metadata.InputModalities) + require.Equal(t, int64(256_000), metadata.ContextWindow) +} + +// Scenario: 上游 /models 增删型号后,正式同步用最新清单替换能力快照。 +func TestSyncUpstreamModelCatalogReplacesSnapshotWhenUpstreamModelsChange(t *testing.T) { + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"data":[ + {"id":"old-model","reasoning":false,"input_modalities":["text"],"context_window":128000}, + {"id":"kept-model","reasoning":false,"input_modalities":["text"],"context_window":128000} + ]}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"data":[ + {"id":"kept-model","reasoning":false,"input_modalities":["text"],"context_window":128000}, + {"id":"new-model","reasoning":false,"input_modalities":["text"],"context_window":256000} + ]}`)), + }, + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + account := &Account{ + ID: 101, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "key", + "base_url": "https://provider.example/v1", + }, + } + + first, err := svc.SyncUpstreamModelCatalog(context.Background(), account) + require.NoError(t, err) + require.Equal(t, []string{"kept-model", "old-model"}, first.Models) + + second, err := svc.SyncUpstreamModelCatalog(context.Background(), account) + require.NoError(t, err) + require.Equal(t, []string{"kept-model", "new-model"}, second.Models) + require.NotContains(t, second.Metadata, "old-model") + require.Contains(t, second.Metadata, "new-model") + + encoded, err := json.Marshal(repo.updates[UpstreamModelMetadataExtraKey]) + require.NoError(t, err) + var snapshot UpstreamModelMetadataSnapshot + require.NoError(t, json.Unmarshal(encoded, &snapshot)) + require.NotContains(t, snapshot.Models, "old-model") + require.Contains(t, snapshot.Models, "new-model") +} + +// Scenario: 上游明确声明无推理能力时保存 false。 +func TestSyncUpstreamModelCatalogPersistsExplicitNonReasoningCapability(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"models":[{ + "id":"company-coding-model", + "display_name":"Company Coding Model", + "reasoning":false, + "input_modalities":["text"], + "context_window":64000 + }]}`)), + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 94, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + }) + require.NoError(t, err) + require.Len(t, upstream.requests, 1) + metadata := catalog.Metadata["company-coding-model"] + require.NotNil(t, metadata.Reasoning) + require.False(t, *metadata.Reasoning) + require.Empty(t, metadata.SupportedReasoningLevels) + require.Equal(t, []string{"text"}, metadata.InputModalities) + require.Equal(t, int64(64_000), metadata.ContextWindow) + require.NotNil(t, repo.updates) +} + +func TestSyncUpstreamModelCatalogClassifiesSnapshotPersistenceFailureAsInternal(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"models":[{ + "id":"company-coding-model", + "reasoning":false, + "input_modalities":["text"], + "context_window":64000 + }]}`)), + }} + repo := &upstreamModelMetadataRepoStub{err: errors.New("database unavailable")} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + _, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 95, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + }) + require.Error(t, err) + var syncErr *UpstreamModelSyncError + require.ErrorAs(t, err, &syncErr) + require.Equal(t, UpstreamModelSyncErrorInternal, syncErr.Kind) +} + +// Scenario: 元数据源失败时保留已有快照。 +func TestSyncUpstreamModelCatalogDoesNotOverwriteSnapshotWhenRegistryFails(t *testing.T) { + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"x-preview-f-free"}]}`))}, + {StatusCode: http.StatusBadGateway, Body: io.NopCloser(strings.NewReader(`{"error":"unavailable"}`))}, + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 93, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://opencode.ai/zen/v1"}, + Extra: map[string]any{UpstreamModelMetadataExtraKey: map[string]any{ + "source": "models.dev", "models": map[string]any{"x-preview-f-free": map[string]any{"reasoning": true}}, + }}, + }) + require.NoError(t, err) + require.Equal(t, []string{"x-preview-f-free"}, catalog.Models) + require.Empty(t, catalog.Metadata) + require.Equal(t, []UpstreamModelSyncWarning{{ + Code: UpstreamModelMetadataIncompleteCode, + Message: "Model IDs were synced, but capability metadata is incomplete.", + }}, catalog.Warnings) + require.Nil(t, repo.updates, "a failed metadata enrichment must not erase a previously saved snapshot") +} + +func TestSyncUpstreamModelCatalogDoesNotPersistPartialMetadataWhenRegistryFails(t *testing.T) { + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"models":[{ + "id":"partially-described-model", + "display_name":"Partial Model" + }]}`))}, + {StatusCode: http.StatusBadGateway, Body: io.NopCloser(strings.NewReader(`{"error":"unavailable"}`))}, + }} + repo := &upstreamModelMetadataRepoStub{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: upstreamModelSyncTestConfig()} + + catalog, err := svc.SyncUpstreamModelCatalog(context.Background(), &Account{ + ID: 96, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"}, + Extra: map[string]any{UpstreamModelMetadataExtraKey: map[string]any{ + "source": "upstream", "models": map[string]any{"partially-described-model": map[string]any{ + "reasoning": true, "supported_reasoning_levels": []any{"low", "high"}, + }}, + }}, + }) + require.NoError(t, err) + require.Equal(t, []string{"partially-described-model"}, catalog.Models) + require.Equal(t, "Partial Model", catalog.Metadata["partially-described-model"].DisplayName) + require.Equal(t, UpstreamModelMetadataIncompleteCode, catalog.Warnings[0].Code) + require.Nil(t, repo.updates, "partial metadata must not replace a more complete persisted snapshot") +} + func TestFetchUpstreamSupportedModelsUsesConfiguredBodyLimit(t *testing.T) { t.Parallel() diff --git a/frontend/src/api/__tests__/codex.spec.ts b/frontend/src/api/__tests__/codex.spec.ts new file mode 100644 index 0000000000..57b19d7557 --- /dev/null +++ b/frontend/src/api/__tests__/codex.spec.ts @@ -0,0 +1,79 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' +import { + buildCodexModelsManifestUrl, + fetchCodexModelsManifest +} from '../codex' + +describe('Codex models API', () => { + afterEach(() => { + vi.unstubAllGlobals() + }) + + it('builds the authenticated Codex manifest endpoint from the public API base', () => { + expect(buildCodexModelsManifestUrl('https://example.com/api/v1/')).toBe( + 'https://example.com/api/v1/models?client_version=0.147.0' + ) + }) + + it('fetches a manifest with the current API key without adding it to the catalog', async () => { + const manifest = { + models: [ + { + slug: 'grok-4.6', + default_reasoning_level: 'high', + supported_reasoning_levels: [ + { effort: 'low', description: 'Fast responses' }, + { effort: 'xhigh', description: 'Extra-high reasoning depth' } + ], + input_modalities: ['text', 'image'], + model_messages: { instructions_template: 'Use the routed model.' } + }, + { + slug: 'deepseek-v4-pro', + default_reasoning_level: 'high', + supported_reasoning_levels: [ + { effort: 'low', description: 'Fast responses' }, + { effort: 'max', description: 'Maximum reasoning depth' } + ], + input_modalities: ['text'], + model_messages: { instructions_template: 'Use the routed model.' } + } + ] + } + const fetchMock = vi.fn().mockResolvedValue({ + ok: true, + status: 200, + json: async () => manifest + }) + vi.stubGlobal('fetch', fetchMock) + + const result = await fetchCodexModelsManifest('https://example.com/v1', 'sk-user-test') + + expect(fetchMock).toHaveBeenCalledWith( + 'https://example.com/v1/models?client_version=0.147.0', + expect.objectContaining({ + headers: { + Accept: 'application/json', + Authorization: 'Bearer sk-user-test' + } + }) + ) + expect(result.modelCount).toBe(2) + expect(JSON.parse(result.content)).toEqual(manifest) + expect(result.content).toContain('"effort": "xhigh"') + expect(result.content).toContain('"input_modalities"') + expect(result.content).toContain('"instructions_template"') + expect(result.content).not.toContain('sk-user-test') + }) + + it('rejects a successful response that is not a Codex manifest', async () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ + ok: true, + status: 200, + json: async () => ({ object: 'list', data: [] }) + })) + + await expect(fetchCodexModelsManifest('https://example.com/v1', 'sk-user-test')) + .rejects.toThrow('valid manifest') + }) +}) diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index 73bfd29af3..ee053ad9ec 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -541,6 +541,25 @@ export async function getAvailableModels(id: number): Promise { export interface SyncUpstreamModelsResult { models: string[] + metadata?: Record + warnings?: UpstreamModelSyncWarning[] +} + +export interface UpstreamModelSyncWarning { + code: string + message: string +} + +export interface UpstreamModelMetadata { + id: string + display_name?: string + description?: string + reasoning?: boolean + default_reasoning_level?: string + supported_reasoning_levels?: string[] + input_modalities?: string[] + context_window?: number + max_output_tokens?: number } /** @@ -558,6 +577,7 @@ export interface SyncUpstreamPreviewParams { type: string base_url?: string api_key: string + model_mapping?: Record } /** diff --git a/frontend/src/api/codex.ts b/frontend/src/api/codex.ts new file mode 100644 index 0000000000..948fa1eee9 --- /dev/null +++ b/frontend/src/api/codex.ts @@ -0,0 +1,56 @@ +export interface CodexModelsManifestResult { + content: string + modelCount: number +} + +const DEFAULT_CODEX_CLIENT_VERSION = '0.147.0' + +function normalizeCodexBaseUrl(baseUrl: string): string { + const fallback = typeof window !== 'undefined' ? window.location.origin : '' + const value = (baseUrl || fallback).trim().replace(/\/+$/, '') + if (!value) return '/v1' + return /\/v1$/i.test(value) ? value : `${value}/v1` +} + +export function buildCodexModelsManifestUrl( + baseUrl: string, + clientVersion = DEFAULT_CODEX_CLIENT_VERSION +): string { + const url = normalizeCodexBaseUrl(baseUrl) + const params = new URLSearchParams({ client_version: clientVersion }) + return `${url}/models?${params.toString()}` +} + +function isCodexModelsManifest(value: unknown): value is { models: unknown[] } { + return typeof value === 'object' && value !== null && Array.isArray((value as { models?: unknown }).models) +} + +export async function fetchCodexModelsManifest( + baseUrl: string, + apiKey: string, + signal?: AbortSignal +): Promise { + const response = await fetch(buildCodexModelsManifestUrl(baseUrl), { + method: 'GET', + headers: { + Accept: 'application/json', + Authorization: `Bearer ${apiKey}` + }, + cache: 'no-store', + signal + }) + + if (!response.ok) { + throw new Error(`Codex models request failed with status ${response.status}`) + } + + const payload: unknown = await response.json() + if (!isCodexModelsManifest(payload)) { + throw new Error('Codex models response is not a valid manifest') + } + + return { + content: JSON.stringify(payload, null, 2), + modelCount: payload.models.length + } +} diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 2c9621f17f..ba30abfdac 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -1402,7 +1402,12 @@
- +

{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }} {{ @@ -1884,7 +1889,12 @@

- +

{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }} {{ t('admin.accounts.supportsAllModels') }} @@ -2220,7 +2230,12 @@

- +

{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }} {{ @@ -4072,11 +4087,17 @@ const syncPreviewCredentials = computed(() => { const baseUrl = isCNPlatform.value && apiProtocol.value === 'adaptive' ? adaptiveBaseUrls.value.chat_completions.trim() || apiKeyBaseUrl.value.trim() : apiKeyBaseUrl.value.trim() + const modelMapping = buildModelMappingObject( + modelRestrictionMode.value, + allowedModels.value, + modelMappings.value + ) return { platform: form.platform, type: form.type, base_url: baseUrl || undefined, - api_key: apiKeyValue.value + api_key: apiKeyValue.value, + ...(modelMapping ? { model_mapping: modelMapping } : {}) } }) @@ -4093,6 +4114,7 @@ const modelMappings = ref([]) const openAICompactModelMappings = ref([]) const modelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist') const allowedModels = ref([]) +const upstreamModelsPreviewed = ref(false) const DEFAULT_POOL_MODE_RETRY_COUNT = 3 const MAX_POOL_MODE_RETRY_COUNT = 10 const DEFAULT_POOL_MODE_RETRY_STATUS_CODES = [401, 403, 429] @@ -4579,6 +4601,7 @@ watch( } // Clear model-related settings allowedModels.value = [] + upstreamModelsPreviewed.value = false modelMappings.value = [] // Antigravity: 默认使用映射模式并填充默认映射 if (newPlatform === 'antigravity') { @@ -4970,6 +4993,23 @@ const submitCreateAccount = async (payload: CreateAccountRequest) => { submitting.value = true try { const account = await adminAPI.accounts.create(withAntigravityConfirmFlag(payload)) + const modelMapping = payload.credentials.model_mapping + const hasConcreteMappedTarget = payload.type === 'apikey' && + typeof modelMapping === 'object' && + modelMapping !== null && + Object.values(modelMapping).some((target) => + typeof target === 'string' && target.trim() !== '' && !target.includes('*') + ) + if (upstreamModelsPreviewed.value || hasConcreteMappedTarget) { + try { + const result = await adminAPI.accounts.syncUpstreamModels(account.id) + if (result.warnings?.some(warning => warning.code === 'upstream_model_metadata_incomplete')) { + appStore.showWarning(t('admin.accounts.syncUpstreamModelsMetadataIncomplete')) + } + } catch { + appStore.showWarning(t('admin.accounts.syncUpstreamModelsFailed')) + } + } if ( payload.type === 'apikey' && payload.upstream_billing_probe_enabled === true @@ -5110,6 +5150,7 @@ const resetForm = () => { grokOAuth.resetState() oauthFlowRef.value?.reset() antigravityMixedChannelConfirmed.value = false + upstreamModelsPreviewed.value = false clearMixedChannelDialog() } diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index e93d4726e1..b1a6d10e97 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -4132,6 +4132,10 @@ const syncAntigravityUpstreamModels = async () => { } } + if (result.warnings?.some((warning) => warning.code === 'upstream_model_metadata_incomplete')) { + appStore.showWarning(t('admin.accounts.syncUpstreamModelsMetadataIncomplete')) + return + } if (addedCount > 0) { appStore.showSuccess(t('admin.accounts.syncUpstreamModelsSuccess', { count: addedCount, total: upstreamModels.length })) } else { diff --git a/frontend/src/components/account/ModelWhitelistSelector.vue b/frontend/src/components/account/ModelWhitelistSelector.vue index b6086c6d6b..a3b1309fbf 100644 --- a/frontend/src/components/account/ModelWhitelistSelector.vue +++ b/frontend/src/components/account/ModelWhitelistSelector.vue @@ -172,6 +172,7 @@ const props = defineProps<{ const emit = defineEmits<{ 'update:modelValue': [value: string[]] + 'upstream-synced': [] }>() const appStore = useAppStore() @@ -312,6 +313,10 @@ const syncUpstreamModels = async () => { return } + if (!props.accountId) { + emit('upstream-synced') + } + const newModels = [...props.modelValue] let addedCount = 0 for (const model of upstreamModels) { @@ -322,6 +327,10 @@ const syncUpstreamModels = async () => { } emit('update:modelValue', newModels) + if (result.warnings?.some(warning => warning.code === 'upstream_model_metadata_incomplete')) { + appStore.showWarning(t('admin.accounts.syncUpstreamModelsMetadataIncomplete')) + return + } if (addedCount > 0) { appStore.showSuccess(t('admin.accounts.syncUpstreamModelsSuccess', { count: addedCount, total: upstreamModels.length })) } else { diff --git a/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts index e5a0c8c9c1..26305c4233 100644 --- a/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts @@ -5,12 +5,16 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' const { createAccountMock, probeUpstreamBillingMock, + syncUpstreamModelsMock, + showWarningMock, importCodexSessionMock, createOpenAICodexPATMock, authIsSimpleMode, } = vi.hoisted(() => ({ createAccountMock: vi.fn(), probeUpstreamBillingMock: vi.fn(), + syncUpstreamModelsMock: vi.fn(), + showWarningMock: vi.fn(), importCodexSessionMock: vi.fn(), createOpenAICodexPATMock: vi.fn(), authIsSimpleMode: { value: true }, @@ -20,7 +24,7 @@ vi.mock('@/stores/app', () => ({ useAppStore: () => ({ showError: vi.fn(), showSuccess: vi.fn(), - showWarning: vi.fn(), + showWarning: showWarningMock, }), })) @@ -37,6 +41,7 @@ vi.mock('@/api/admin', () => ({ accounts: { create: createAccountMock, probeUpstreamBilling: probeUpstreamBillingMock, + syncUpstreamModels: syncUpstreamModelsMock, checkMixedChannelRisk: vi.fn().mockResolvedValue({ has_risk: false }), importCodexSession: importCodexSessionMock, createOpenAICodexPAT: createOpenAICodexPATMock, @@ -120,8 +125,12 @@ const ModelWhitelistSelectorStub = defineComponent({ platform: String, syncCredentials: Object, }, - emits: ['update:modelValue'], - template: '

', + emits: ['update:modelValue', 'upstream-synced'], + template: ``, }) function mountModal(groups: any[] = []) { @@ -190,6 +199,8 @@ describe('CreateAccountModal OpenAI long-context billing', () => { authIsSimpleMode.value = true createAccountMock.mockReset().mockResolvedValue({ id: 42, platform: 'openai', type: 'apikey' }) probeUpstreamBillingMock.mockReset().mockResolvedValue({}) + syncUpstreamModelsMock.mockReset().mockResolvedValue({ models: [], metadata: {} }) + showWarningMock.mockReset() importCodexSessionMock.mockReset().mockResolvedValue({ created: 1, updated: 0, @@ -236,6 +247,71 @@ describe('CreateAccountModal OpenAI long-context billing', () => { expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) }) + it('persists upstream model metadata after creating an account from preview', async () => { + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + await selectButtonByText(wrapper, 'API Key') + await wrapper.get('form#create-account-form input[type="text"]').setValue('OpenCode account') + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + await wrapper.get('[data-testid="model-whitelist-selector"]').trigger('click') + await wrapper.get('form#create-account-form').trigger('submit.prevent') + await flushPromises() + + expect(createAccountMock).toHaveBeenCalledOnce() + expect(syncUpstreamModelsMock).toHaveBeenCalledWith(42) + }) + + it('includes the current concrete model mapping in preview credentials', async () => { + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + await selectButtonByText(wrapper, 'API Key') + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + await wrapper.get('[data-testid="model-whitelist-selector"]').trigger('click') + await flushPromises() + + expect(wrapper.getComponent(ModelWhitelistSelectorStub).props('syncCredentials')).toMatchObject({ + model_mapping: { 'public-glm': 'public-glm' } + }) + }) + + it('runs formal capability sync after creating an account with explicit mappings', async () => { + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + await selectButtonByText(wrapper, 'API Key') + await wrapper.get('form#create-account-form input[type="text"]').setValue('Mapped account') + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + await selectButtonByText(wrapper, 'admin.accounts.modelMapping') + await selectButtonByText(wrapper, 'admin.accounts.addMapping') + await wrapper.get('input[placeholder="admin.accounts.requestModel"]').setValue('public-glm') + await wrapper.get('input[placeholder="admin.accounts.actualModel"]').setValue('glm-5.3') + await wrapper.get('form#create-account-form').trigger('submit.prevent') + await flushPromises() + + expect(createAccountMock.mock.calls[0]?.[0]?.credentials?.model_mapping).toEqual({ + 'public-glm': 'glm-5.3' + }) + expect(syncUpstreamModelsMock).toHaveBeenCalledWith(42) + }) + + it('warns when post-create capability metadata remains incomplete', async () => { + syncUpstreamModelsMock.mockResolvedValue({ + models: ['x-preview-f-free'], + warnings: [{ code: 'upstream_model_metadata_incomplete', message: 'metadata incomplete' }], + }) + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + await selectButtonByText(wrapper, 'API Key') + await wrapper.get('form#create-account-form input[type="text"]').setValue('OpenCode account') + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + await wrapper.get('[data-testid="model-whitelist-selector"]').trigger('click') + await wrapper.get('form#create-account-form').trigger('submit.prevent') + await flushPromises() + + expect(showWarningMock).toHaveBeenCalledWith( + 'admin.accounts.syncUpstreamModelsMetadataIncomplete' + ) + }) + // namespace 摊平是仅 OAuth 的兼容开关:API Key 走 chat completions 回退桥时由桥自行摊平 it('shows the Codex namespace flatten toggle only for OpenAI OAuth accounts', async () => { const wrapper = mountModal() diff --git a/frontend/src/components/account/__tests__/ModelWhitelistSelector.spec.ts b/frontend/src/components/account/__tests__/ModelWhitelistSelector.spec.ts index 51f18611cc..928ab6e721 100644 --- a/frontend/src/components/account/__tests__/ModelWhitelistSelector.spec.ts +++ b/frontend/src/components/account/__tests__/ModelWhitelistSelector.spec.ts @@ -1,7 +1,23 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' import { flushPromises, mount } from '@vue/test-utils' -const copyToClipboard = vi.fn().mockResolvedValue(true) +const { + copyToClipboard, + showError, + showSuccess, + showInfo, + showWarning, + syncUpstreamModels, + syncUpstreamModelsPreview +} = vi.hoisted(() => ({ + copyToClipboard: vi.fn().mockResolvedValue(true), + showError: vi.fn(), + showSuccess: vi.fn(), + showInfo: vi.fn(), + showWarning: vi.fn(), + syncUpstreamModels: vi.fn(), + syncUpstreamModelsPreview: vi.fn() +})) vi.mock('vue-i18n', async () => { const actual = await vi.importActual('vue-i18n') @@ -15,12 +31,20 @@ vi.mock('vue-i18n', async () => { vi.mock('@/stores/app', () => ({ useAppStore: () => ({ - showError: vi.fn(), - showSuccess: vi.fn(), - showInfo: vi.fn() + showError, + showSuccess, + showInfo, + showWarning }) })) +vi.mock('@/api/admin/accounts', () => ({ + accountsAPI: { + syncUpstreamModels, + syncUpstreamModelsPreview + } +})) + vi.mock('@/composables/useClipboard', () => ({ useClipboard: () => ({ copyToClipboard @@ -29,11 +53,12 @@ vi.mock('@/composables/useClipboard', () => ({ import ModelWhitelistSelector from '../ModelWhitelistSelector.vue' -function mountSelector() { +function mountSelector(props: Record = {}) { return mount(ModelWhitelistSelector, { props: { modelValue: [], - platform: 'openai' + platform: 'openai', + ...props, }, global: { stubs: { @@ -58,6 +83,12 @@ function findModelRow(wrapper: ReturnType, modelId: string describe('ModelWhitelistSelector', () => { beforeEach(() => { copyToClipboard.mockClear() + showError.mockReset() + showSuccess.mockReset() + showInfo.mockReset() + showWarning.mockReset() + syncUpstreamModels.mockReset() + syncUpstreamModelsPreview.mockReset() }) it('copies a model ID without selecting the model', async () => { @@ -86,4 +117,71 @@ describe('ModelWhitelistSelector', () => { expect(wrapper.emitted('update:modelValue')).toEqual([[['gpt-5.6-sol']]]) expect(copyToClipboard).not.toHaveBeenCalled() }) + + it('warns when model IDs sync but capability metadata is incomplete', async () => { + syncUpstreamModels.mockResolvedValue({ + models: ['x-preview-f-free'], + warnings: [ + { + code: 'upstream_model_metadata_incomplete', + message: 'Model IDs were synced, but capability metadata could not be updated.' + } + ] + }) + const wrapper = mount(ModelWhitelistSelector, { + props: { + modelValue: [], + platform: 'openai', + accountId: 46 + }, + global: { + stubs: { + ModelIcon: true + } + } + }) + + const syncButton = wrapper + .findAll('button') + .find(button => button.text() === 'admin.accounts.syncUpstreamModels') + expect(syncButton).toBeDefined() + await syncButton!.trigger('click') + await flushPromises() + + expect(wrapper.emitted('update:modelValue')).toEqual([[['x-preview-f-free']]]) + expect(showWarning).toHaveBeenCalledWith('admin.accounts.syncUpstreamModelsMetadataIncomplete') + expect(showSuccess).not.toHaveBeenCalled() + }) + + it('reports a successful preview so account creation can persist metadata', async () => { + syncUpstreamModelsPreview.mockResolvedValue({ + models: ['x-preview-f-free'], + metadata: { + 'x-preview-f-free': { + id: 'x-preview-f-free', + reasoning: true, + supported_reasoning_levels: ['low', 'high', 'max'], + }, + }, + }) + const wrapper = mountSelector({ + syncCredentials: { + platform: 'openai', + type: 'apikey', + base_url: 'https://opencode.ai/zen/v1', + api_key: 'test-key', + }, + }) + const syncButton = wrapper + .findAll('button') + .find(button => button.text() === 'admin.accounts.syncUpstreamModels') + + expect(syncButton).toBeDefined() + await syncButton?.trigger('click') + await flushPromises() + + expect(syncUpstreamModelsPreview).toHaveBeenCalledOnce() + expect(wrapper.emitted('upstream-synced')).toEqual([[]]) + expect(wrapper.emitted('update:modelValue')).toEqual([[['x-preview-f-free']]]) + }) }) diff --git a/frontend/src/components/keys/UseKeyModal.vue b/frontend/src/components/keys/UseKeyModal.vue index c8ae79d229..fc5356cf5e 100644 --- a/frontend/src/components/keys/UseKeyModal.vue +++ b/frontend/src/components/keys/UseKeyModal.vue @@ -172,6 +172,65 @@
+
+
+
+

+ {{ t('keys.useKeyModal.codexModelCatalog.title') }} +

+

+ {{ t('keys.useKeyModal.codexModelCatalog.description') }} +

+

+ {{ codexModelCatalogPath }} +

+
+ + +
+

+ {{ t('keys.useKeyModal.codexModelCatalog.modelsCount', { count: codexModelManifestModelCount }) }} +

+

+ {{ t('keys.useKeyModal.codexModelCatalog.errorDescription') }} +

+
+
@@ -198,10 +257,18 @@