mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-06 16:03:49 +08:00
Merge branch 'main' into fix/openai-reasoning-replay
This commit is contained in:
@@ -58,6 +58,11 @@ Please read the following carefully before using this project:
|
||||
<td>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 <a href="https://aigocode.com/invite/SUB2API">this link</a>, you'll receive an extra 10% bonus credit on your first top-up!</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
|
||||
<td>Real GPT-5.6 series at 3% of OpenAI pricing — <a href="https://codex-everywhere.com">CodexEverywhere</a> 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 <a href="https://codex-everywhere.com">codex-everywhere.com</a>.</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
|
||||
<td>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 <a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus - Premium AI Accounts & Top-ups</a>, users can unlock the mind-blowing rate of 10% of the official GPT subscription price (90% OFF)</td>
|
||||
|
||||
@@ -59,6 +59,11 @@
|
||||
<td>感谢 AIGoCode 赞助了本项目!AIGoCode 是一站式集成 Claude Code、Codex 以及最新 Gemini 模型的综合平台,为您提供稳定、高效、高性价比的 AI 编程服务。平台提供灵活的订阅方案,零封号风险,免 VPN 直连,响应极速。AIGoCode 为 sub2api 用户准备了专属福利:通过<a href="https://aigocode.com/invite/SUB2API">此链接</a>注册,首次充值可额外获得 10% 赠送额度!</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
|
||||
<td>Real GPT-5.6 series at 3% of OpenAI pricing — <a href="https://codex-everywhere.com">CodexEverywhere</a> 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 <a href="https://codex-everywhere.com">codex-everywhere.com</a>.</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
|
||||
<td>感谢 BmoPlus 赞助了本项目!BmoPlus 是一家专为AI订阅重度用户打造的可靠 AI 账号代充服务商,提供稳定的 ChatGPT Plus / ChatGPT Pro(全程质保) / Claude Pro / Super Grok / Gemini Pro 的官方代充&成品账号。 通过<a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus AI成品号专卖/代充</a>注册下单的用户,可享GPT 官网订阅一折 的震撼价格!</td>
|
||||
|
||||
@@ -58,6 +58,11 @@
|
||||
<td>AIGoCode のご支援に感謝します!AIGoCode は Claude Code、Codex、最新の Gemini モデルを統合したオールインワンプラットフォームで、安定的かつ効率的でコストパフォーマンスに優れた AI コーディングサービスを提供します。柔軟なサブスクリプションプラン、アカウント停止リスクゼロ、VPN 不要の直接アクセス、超高速レスポンスが特長です。AIGoCode は sub2api ユーザー向けに特別特典を用意しています:<a href="https://aigocode.com/invite/SUB2API">こちらのリンク</a>から登録すると、初回チャージ時に 10% のボーナスクレジットを追加プレゼント!</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
|
||||
<td>OpenAI 公式価格のわずか 3% で本物の GPT-5.6 シリーズを提供 — <a href="https://codex-everywhere.com">CodexEverywhere</a> は世界中の開発者にフロンティアモデルへのアクセスを民主化しています。私たちは透明性と誠実さを信条とし、モデル品質は数か月にわたるアクティブなコミュニティの監視によって検証されています。USD および暗号通貨に対応。<a href="https://codex-everywhere.com">codex-everywhere.com</a> で $20 の無料トライアルから始めましょう。</td>
|
||||
</tr>
|
||||
|
||||
<tr>
|
||||
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
|
||||
<td>本プロジェクトにご支援いただいた BmoPlus に感謝いたします!BmoPlusは、AIサブスクリプションのヘビーユーザー向けに特化した信頼性の高いAIアカウントサービスプロバイダーであり、安定した ChatGPT Plus / ChatGPT Pro (完全保証) / Claude Pro / Super Grok / Gemini Pro の公式代行チャージおよび即納アカウントを提供しています。こちらの<a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus AIアカウント専門店/代行チャージ</a>経由でご登録・ご注文いただいたユーザー様は、GPTを 公式サイト価格の約1割(90% OFF) という驚異的な価格でご利用いただけます!</td>
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 5.9 KiB |
@@ -1 +1 @@
|
||||
0.1.182
|
||||
0.1.183
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 != "" {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = ""
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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])
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 渲染下面的表达式:去掉存量邮箱的大小写、首尾空白(与
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 初始化网关调试日志文件。
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"])
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()}})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 协议账号原样返回。
|
||||
|
||||
@@ -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}`)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
"文字兜底判据是可重试的——正因如此,写账号冷却会沿号池级联")
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 只做历史摘要、不需要模型寻址工具,
|
||||
// 因此保持既有的摊平 + 全量清理行为,不随默认值翻转扩大风险面。
|
||||
|
||||
@@ -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,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 != "" {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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')
|
||||
})
|
||||
})
|
||||
@@ -541,6 +541,25 @@ export async function getAvailableModels(id: number): Promise<ClaudeModel[]> {
|
||||
|
||||
export interface SyncUpstreamModelsResult {
|
||||
models: string[]
|
||||
metadata?: Record<string, UpstreamModelMetadata>
|
||||
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<string, string>
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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<CodexModelsManifestResult> {
|
||||
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
|
||||
}
|
||||
}
|
||||
@@ -1402,7 +1402,12 @@
|
||||
|
||||
<!-- Whitelist Mode -->
|
||||
<div v-if="modelRestrictionMode === 'whitelist'">
|
||||
<ModelWhitelistSelector v-model="allowedModels" :platform="form.platform" :sync-credentials="syncPreviewCredentials" />
|
||||
<ModelWhitelistSelector
|
||||
v-model="allowedModels"
|
||||
:platform="form.platform"
|
||||
:sync-credentials="syncPreviewCredentials"
|
||||
@upstream-synced="upstreamModelsPreviewed = true"
|
||||
/>
|
||||
<p class="text-xs text-gray-500 dark:text-gray-400">
|
||||
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
|
||||
<span v-if="allowedModels.length === 0">{{
|
||||
@@ -1884,7 +1889,12 @@
|
||||
|
||||
<!-- Whitelist Mode -->
|
||||
<div v-if="modelRestrictionMode === 'whitelist'">
|
||||
<ModelWhitelistSelector v-model="allowedModels" platform="anthropic" :sync-credentials="syncPreviewCredentials" />
|
||||
<ModelWhitelistSelector
|
||||
v-model="allowedModels"
|
||||
platform="anthropic"
|
||||
:sync-credentials="syncPreviewCredentials"
|
||||
@upstream-synced="upstreamModelsPreviewed = true"
|
||||
/>
|
||||
<p class="text-xs text-gray-500 dark:text-gray-400">
|
||||
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
|
||||
<span v-if="allowedModels.length === 0">{{ t('admin.accounts.supportsAllModels') }}</span>
|
||||
@@ -2220,7 +2230,12 @@
|
||||
|
||||
<!-- Whitelist Mode -->
|
||||
<div v-if="modelRestrictionMode === 'whitelist'">
|
||||
<ModelWhitelistSelector v-model="allowedModels" :platform="form.platform" :sync-credentials="syncPreviewCredentials" />
|
||||
<ModelWhitelistSelector
|
||||
v-model="allowedModels"
|
||||
:platform="form.platform"
|
||||
:sync-credentials="syncPreviewCredentials"
|
||||
@upstream-synced="upstreamModelsPreviewed = true"
|
||||
/>
|
||||
<p class="text-xs text-gray-500 dark:text-gray-400">
|
||||
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
|
||||
<span v-if="allowedModels.length === 0">{{
|
||||
@@ -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<ModelMapping[]>([])
|
||||
const openAICompactModelMappings = ref<ModelMapping[]>([])
|
||||
const modelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist')
|
||||
const allowedModels = ref<string[]>([])
|
||||
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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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: '<div data-testid="model-whitelist-selector" />',
|
||||
emits: ['update:modelValue', 'upstream-synced'],
|
||||
template: `<button
|
||||
type="button"
|
||||
data-testid="model-whitelist-selector"
|
||||
@click="$emit('update:modelValue', ['public-glm']); $emit('upstream-synced')"
|
||||
>models</button>`,
|
||||
})
|
||||
|
||||
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()
|
||||
|
||||
@@ -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<typeof import('vue-i18n')>('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<string, unknown> = {}) {
|
||||
return mount(ModelWhitelistSelector, {
|
||||
props: {
|
||||
modelValue: [],
|
||||
platform: 'openai'
|
||||
platform: 'openai',
|
||||
...props,
|
||||
},
|
||||
global: {
|
||||
stubs: {
|
||||
@@ -58,6 +83,12 @@ function findModelRow(wrapper: ReturnType<typeof mountSelector>, 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']]])
|
||||
})
|
||||
})
|
||||
|
||||
@@ -172,6 +172,65 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<section
|
||||
v-if="showCodexModelCatalog"
|
||||
data-testid="codex-model-catalog"
|
||||
class="overflow-hidden rounded-lg border border-gray-200 bg-gray-50 dark:border-dark-700 dark:bg-dark-800/50"
|
||||
>
|
||||
<div class="flex flex-col gap-3 px-4 py-3 sm:flex-row sm:items-center sm:justify-between">
|
||||
<div class="min-w-0">
|
||||
<h3 class="text-sm font-medium text-gray-900 dark:text-white">
|
||||
{{ t('keys.useKeyModal.codexModelCatalog.title') }}
|
||||
</h3>
|
||||
<p class="mt-1 text-xs text-gray-500 dark:text-gray-400">
|
||||
{{ t('keys.useKeyModal.codexModelCatalog.description') }}
|
||||
</p>
|
||||
<p class="mt-1 truncate font-mono text-xs text-gray-700 dark:text-gray-300">
|
||||
{{ codexModelCatalogPath }}
|
||||
</p>
|
||||
</div>
|
||||
<button
|
||||
v-if="codexModelManifestState === 'ready'"
|
||||
type="button"
|
||||
class="btn btn-primary min-h-9 flex-shrink-0 px-3 text-xs"
|
||||
@click="downloadCodexModelManifest"
|
||||
>
|
||||
<Icon name="download" size="sm" class="mr-1.5" />
|
||||
{{ t('keys.useKeyModal.codexModelCatalog.download') }}
|
||||
</button>
|
||||
<button
|
||||
v-else
|
||||
type="button"
|
||||
data-testid="codex-model-catalog-fetch"
|
||||
class="btn btn-primary min-h-9 flex-shrink-0 px-3 text-xs"
|
||||
:disabled="codexModelManifestState === 'loading' || !apiKey"
|
||||
@click="loadCodexModelManifest"
|
||||
>
|
||||
<Icon
|
||||
name="refresh"
|
||||
size="sm"
|
||||
class="mr-1.5"
|
||||
:class="codexModelManifestState === 'loading' ? 'animate-spin' : ''"
|
||||
/>
|
||||
{{ codexModelManifestState === 'error'
|
||||
? t('keys.useKeyModal.codexModelCatalog.retry')
|
||||
: t('keys.useKeyModal.codexModelCatalog.fetch') }}
|
||||
</button>
|
||||
</div>
|
||||
<p
|
||||
v-if="codexModelManifestState === 'ready'"
|
||||
class="border-t border-gray-200 px-4 py-2 text-xs text-emerald-700 dark:border-dark-700 dark:text-emerald-300"
|
||||
>
|
||||
{{ t('keys.useKeyModal.codexModelCatalog.modelsCount', { count: codexModelManifestModelCount }) }}
|
||||
</p>
|
||||
<p
|
||||
v-else-if="codexModelManifestState === 'error'"
|
||||
class="border-t border-red-200 px-4 py-2 text-xs text-red-700 dark:border-red-900 dark:text-red-300"
|
||||
>
|
||||
{{ t('keys.useKeyModal.codexModelCatalog.errorDescription') }}
|
||||
</p>
|
||||
</section>
|
||||
|
||||
<!-- Usage Note -->
|
||||
<div v-if="showPlatformNote" class="flex items-start gap-3 p-3 rounded-lg bg-blue-50 dark:bg-blue-900/20 border border-blue-100 dark:border-blue-800">
|
||||
<Icon name="infoCircle" size="md" class="text-blue-500 flex-shrink-0 mt-0.5" />
|
||||
@@ -198,10 +257,18 @@
|
||||
<script setup lang="ts">
|
||||
import { ref, computed, h, watch, type Component } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import { saveAs } from 'file-saver'
|
||||
import BaseDialog from '@/components/common/BaseDialog.vue'
|
||||
import Icon from '@/components/icons/Icon.vue'
|
||||
import { useClipboard } from '@/composables/useClipboard'
|
||||
import { fetchCodexModelsManifest } from '@/api/codex'
|
||||
import type { GroupPlatform } from '@/types'
|
||||
import {
|
||||
findCodexCatalogModel,
|
||||
formatCodexReasoningEffortTomlLine,
|
||||
parseCodexCatalogModels,
|
||||
selectCodexConfigReasoningEffort
|
||||
} from '@/utils/codexCatalogConfig'
|
||||
|
||||
interface Props {
|
||||
show: boolean
|
||||
@@ -239,6 +306,29 @@ const activeTab = ref<string>('unix')
|
||||
const activeClientTab = ref<string>('claude')
|
||||
type CodexAuthMode = 'legacy' | 'api-key'
|
||||
const codexAuthMode = ref<CodexAuthMode>('legacy')
|
||||
type CodexModelManifestState = 'idle' | 'loading' | 'ready' | 'error'
|
||||
const codexModelManifestState = ref<CodexModelManifestState>('idle')
|
||||
const codexModelManifestContent = ref('')
|
||||
const codexModelManifestModelCount = ref(0)
|
||||
let codexModelManifestController: AbortController | null = null
|
||||
let codexModelManifestRequestID = 0
|
||||
|
||||
const showCodexModelCatalog = computed(() =>
|
||||
props.show &&
|
||||
(activeClientTab.value === 'codex' ||
|
||||
(props.platform === 'openai' && activeClientTab.value === 'codex-ws'))
|
||||
)
|
||||
|
||||
const codexModelCatalogPath = computed(() => {
|
||||
const isWindows = activeTab.value === 'windows'
|
||||
const configDir = isWindows ? '%userprofile%\\.codex' : '~/.codex'
|
||||
return joinConfigPath(configDir, 'codex-models.json', isWindows)
|
||||
})
|
||||
|
||||
const codexManifestContext = computed(() => {
|
||||
if (!showCodexModelCatalog.value) return ''
|
||||
return `${props.platform}|${props.baseUrl}|${props.apiKey}`
|
||||
})
|
||||
|
||||
// Reset tabs when platform changes
|
||||
const defaultClientTab = computed(() => {
|
||||
@@ -265,6 +355,14 @@ watch(() => props.platform, () => {
|
||||
watch(() => props.show, (show) => {
|
||||
if (show) {
|
||||
codexAuthMode.value = 'legacy'
|
||||
} else {
|
||||
resetCodexModelManifest()
|
||||
}
|
||||
})
|
||||
|
||||
watch(codexManifestContext, (context, previousContext) => {
|
||||
if (context !== previousContext) {
|
||||
resetCodexModelManifest()
|
||||
}
|
||||
})
|
||||
|
||||
@@ -353,12 +451,14 @@ const clientTabs = computed((): TabConfig[] => {
|
||||
case 'gemini':
|
||||
return [
|
||||
{ id: 'gemini', label: t('keys.useKeyModal.cliTabs.geminiCli'), icon: SparkleIcon },
|
||||
{ id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon },
|
||||
{ id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon }
|
||||
]
|
||||
case 'antigravity':
|
||||
return [
|
||||
{ id: 'claude', label: t('keys.useKeyModal.cliTabs.claudeCode'), icon: TerminalIcon },
|
||||
{ id: 'gemini', label: t('keys.useKeyModal.cliTabs.geminiCli'), icon: SparkleIcon },
|
||||
{ id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon },
|
||||
{ id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon }
|
||||
]
|
||||
case 'grok':
|
||||
@@ -368,9 +468,17 @@ const clientTabs = computed((): TabConfig[] => {
|
||||
{ id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon },
|
||||
{ id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon }
|
||||
]
|
||||
case 'deepseek':
|
||||
case 'composite':
|
||||
return [
|
||||
{ id: 'claude', label: t('keys.useKeyModal.cliTabs.claudeCode'), icon: TerminalIcon },
|
||||
{ id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon },
|
||||
{ id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon }
|
||||
]
|
||||
default:
|
||||
return [
|
||||
{ id: 'claude', label: t('keys.useKeyModal.cliTabs.claudeCode'), icon: TerminalIcon },
|
||||
{ id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon },
|
||||
{ id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon }
|
||||
]
|
||||
}
|
||||
@@ -405,6 +513,13 @@ const currentTabs = computed(() => {
|
||||
})
|
||||
|
||||
const platformDescription = computed(() => {
|
||||
if (activeClientTab.value === 'codex' &&
|
||||
props.platform !== 'openai' &&
|
||||
props.platform !== 'grok' &&
|
||||
props.platform !== 'deepseek' &&
|
||||
props.platform !== 'composite') {
|
||||
return t('keys.useKeyModal.routedCodex.description')
|
||||
}
|
||||
switch (props.platform) {
|
||||
case 'openai':
|
||||
if (activeClientTab.value === 'claude') {
|
||||
@@ -423,12 +538,27 @@ const platformDescription = computed(() => {
|
||||
return t('keys.useKeyModal.grok.codexDescription')
|
||||
}
|
||||
return t('keys.useKeyModal.grok.description')
|
||||
case 'deepseek':
|
||||
return activeClientTab.value === 'codex'
|
||||
? t('keys.useKeyModal.deepseek.codexDescription')
|
||||
: t('keys.useKeyModal.deepseek.description')
|
||||
case 'composite':
|
||||
return activeClientTab.value === 'codex'
|
||||
? t('keys.useKeyModal.composite.codexDescription')
|
||||
: t('keys.useKeyModal.composite.description')
|
||||
default:
|
||||
return t('keys.useKeyModal.description')
|
||||
}
|
||||
})
|
||||
|
||||
const platformNote = computed(() => {
|
||||
if (activeClientTab.value === 'codex' &&
|
||||
props.platform !== 'openai' &&
|
||||
props.platform !== 'grok' &&
|
||||
props.platform !== 'deepseek' &&
|
||||
props.platform !== 'composite') {
|
||||
return t('keys.useKeyModal.routedCodex.note')
|
||||
}
|
||||
switch (props.platform) {
|
||||
case 'openai':
|
||||
if (activeClientTab.value === 'claude') {
|
||||
@@ -460,6 +590,14 @@ const platformNote = computed(() => {
|
||||
return t('keys.useKeyModal.grok.noteWindows')
|
||||
}
|
||||
return t('keys.useKeyModal.grok.note')
|
||||
case 'deepseek':
|
||||
return activeClientTab.value === 'codex'
|
||||
? t('keys.useKeyModal.deepseek.codexNote')
|
||||
: t('keys.useKeyModal.note')
|
||||
case 'composite':
|
||||
return activeClientTab.value === 'codex'
|
||||
? t('keys.useKeyModal.composite.codexNote')
|
||||
: t('keys.useKeyModal.note')
|
||||
default:
|
||||
return t('keys.useKeyModal.note')
|
||||
}
|
||||
@@ -467,6 +605,66 @@ const platformNote = computed(() => {
|
||||
|
||||
const showPlatformNote = computed(() => activeClientTab.value !== 'opencode')
|
||||
|
||||
function resetCodexModelManifest() {
|
||||
codexModelManifestController?.abort()
|
||||
codexModelManifestController = null
|
||||
codexModelManifestRequestID += 1
|
||||
codexModelManifestState.value = 'idle'
|
||||
codexModelManifestContent.value = ''
|
||||
codexModelManifestModelCount.value = 0
|
||||
}
|
||||
|
||||
async function loadCodexModelManifest() {
|
||||
if (!showCodexModelCatalog.value || !props.apiKey) return
|
||||
|
||||
codexModelManifestController?.abort()
|
||||
const controller = new AbortController()
|
||||
const requestID = ++codexModelManifestRequestID
|
||||
codexModelManifestController = controller
|
||||
codexModelManifestState.value = 'loading'
|
||||
|
||||
try {
|
||||
const result = await fetchCodexModelsManifest(props.baseUrl, props.apiKey, controller.signal)
|
||||
if (requestID !== codexModelManifestRequestID) return
|
||||
codexModelManifestContent.value = result.content
|
||||
codexModelManifestModelCount.value = result.modelCount
|
||||
codexModelManifestState.value = 'ready'
|
||||
} catch (error) {
|
||||
const errorName = error && typeof error === 'object' && 'name' in error
|
||||
? String((error as { name?: unknown }).name || '')
|
||||
: ''
|
||||
if (requestID !== codexModelManifestRequestID || errorName === 'AbortError') return
|
||||
codexModelManifestState.value = 'error'
|
||||
} finally {
|
||||
if (requestID === codexModelManifestRequestID) {
|
||||
codexModelManifestController = null
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function downloadCodexModelManifest() {
|
||||
if (!codexModelManifestContent.value) return
|
||||
saveAs(
|
||||
new Blob([codexModelManifestContent.value], { type: 'application/json;charset=utf-8' }),
|
||||
'codex-models.json'
|
||||
)
|
||||
}
|
||||
|
||||
const codexCatalogModelSlugs = computed(() =>
|
||||
parseCodexCatalogModels(codexModelManifestContent.value).map((model) => model.slug)
|
||||
)
|
||||
|
||||
function selectCodexCatalogModel(preferredModel: string): string {
|
||||
if (codexCatalogModelSlugs.value.includes(preferredModel)) return preferredModel
|
||||
return codexCatalogModelSlugs.value[0] || preferredModel
|
||||
}
|
||||
|
||||
function codexReasoningEffortTomlLine(modelSlug: string): string {
|
||||
return formatCodexReasoningEffortTomlLine(
|
||||
selectCodexConfigReasoningEffort(findCodexCatalogModel(codexModelManifestContent.value, modelSlug))
|
||||
)
|
||||
}
|
||||
|
||||
const escapeHtml = (value: string) => value
|
||||
.replace(/&/g, '&')
|
||||
.replace(/</g, '<')
|
||||
@@ -534,8 +732,14 @@ const currentFiles = computed((): FileConfig[] => {
|
||||
}
|
||||
return generateOpenAIFiles(baseUrl, apiKey)
|
||||
case 'gemini':
|
||||
if (activeClientTab.value === 'codex') {
|
||||
return generateRoutedCodexFiles(apiBase, apiKey, 'gemini')
|
||||
}
|
||||
return [generateGeminiCliContent(baseUrl, apiKey)]
|
||||
case 'antigravity':
|
||||
if (activeClientTab.value === 'codex') {
|
||||
return generateRoutedCodexFiles(apiBase, apiKey, 'antigravity')
|
||||
}
|
||||
if (activeClientTab.value === 'gemini') {
|
||||
return [generateGeminiCliContent(`${baseUrl}/antigravity`, apiKey)]
|
||||
}
|
||||
@@ -548,7 +752,20 @@ const currentFiles = computed((): FileConfig[] => {
|
||||
return generateGrokCodexFiles(apiBase, apiKey)
|
||||
}
|
||||
return generateGrokFiles(apiBase, apiKey)
|
||||
case 'deepseek':
|
||||
if (activeClientTab.value === 'codex') {
|
||||
return generateRoutedCodexFiles(apiBase, apiKey, 'deepseek')
|
||||
}
|
||||
return generateAnthropicFiles(baseRoot, apiKey)
|
||||
case 'composite':
|
||||
if (activeClientTab.value === 'codex') {
|
||||
return generateRoutedCodexFiles(apiBase, apiKey, 'composite')
|
||||
}
|
||||
return generateAnthropicFiles(baseRoot, apiKey)
|
||||
default:
|
||||
if (activeClientTab.value === 'codex' && props.platform) {
|
||||
return generateRoutedCodexFiles(apiBase, apiKey, props.platform)
|
||||
}
|
||||
return generateAnthropicFiles(baseUrl, apiKey)
|
||||
}
|
||||
})
|
||||
@@ -714,12 +931,15 @@ function generateOpenAIFiles(baseUrl: string, apiKey: string): FileConfig[] {
|
||||
const isWindows = activeTab.value === 'windows'
|
||||
const configDir = isWindows ? '%userprofile%\\.codex' : '~/.codex'
|
||||
|
||||
const model = selectCodexCatalogModel('gpt-5.5')
|
||||
const reasoningEffortLine = codexReasoningEffortTomlLine(model)
|
||||
|
||||
// config.toml content
|
||||
const configContent = `model_provider = "OpenAI"
|
||||
model = "gpt-5.5"
|
||||
review_model = "gpt-5.5"
|
||||
model_reasoning_effort = "xhigh"
|
||||
disable_response_storage = true
|
||||
model = "${model}"
|
||||
review_model = "${model}"
|
||||
${reasoningEffortLine}disable_response_storage = true
|
||||
model_catalog_json = "${escapeTomlBasicString(codexModelCatalogPath.value)}"
|
||||
network_access = "enabled"
|
||||
windows_wsl_setup_acknowledged = true
|
||||
|
||||
@@ -764,6 +984,10 @@ function joinConfigPath(dir: string, file: string, windows: boolean): string {
|
||||
return `${dir}\\${file}`
|
||||
}
|
||||
|
||||
function escapeTomlBasicString(value: string): string {
|
||||
return value.replace(/\\/g, '\\\\').replace(/"/g, '\\"')
|
||||
}
|
||||
|
||||
function generateGrokFiles(baseUrl: string, apiKey: string): FileConfig[] {
|
||||
// Prefer unix/cmd/powershell when shell tabs are shown; fall back to windows tab.
|
||||
const shell = activeTab.value
|
||||
@@ -912,6 +1136,7 @@ function generateGrokCodexFiles(baseUrl: string, apiKey: string): FileConfig[] {
|
||||
const shell = activeTab.value
|
||||
const isWindowsPath = shell === 'windows' || shell === 'cmd' || shell === 'powershell'
|
||||
const configDir = isWindowsPath ? '%userprofile%\\.codex' : '~/.codex'
|
||||
const model = selectCodexCatalogModel('grok-4.5')
|
||||
|
||||
let envPath: string
|
||||
let envContent: string
|
||||
@@ -937,9 +1162,10 @@ function generateGrokCodexFiles(baseUrl: string, apiKey: string): FileConfig[] {
|
||||
# Switch model: grok-4.5 | grok-4.3 | grok-build-0.1 | grok-4.20-multi-agent-0309 (text / web_search)
|
||||
|
||||
model_provider = "sub2api"
|
||||
model = "grok-4.5"
|
||||
model = "${model}"
|
||||
model_catalog_json = "${escapeTomlBasicString(codexModelCatalogPath.value)}"
|
||||
# Optional:
|
||||
# review_model = "grok-4.5"
|
||||
# review_model = "${model}"
|
||||
# model_reasoning_effort = "medium"
|
||||
# model_context_window = 500000
|
||||
# disable_response_storage = true
|
||||
@@ -973,16 +1199,83 @@ supports_websockets = false
|
||||
]
|
||||
}
|
||||
|
||||
function generateRoutedCodexFiles(
|
||||
baseUrl: string,
|
||||
apiKey: string,
|
||||
platform: GroupPlatform
|
||||
): FileConfig[] {
|
||||
const isWindows = activeTab.value === 'windows'
|
||||
const configDir = isWindows ? '%userprofile%\\.codex' : '~/.codex'
|
||||
const preferredModels: Partial<Record<GroupPlatform, string>> = {
|
||||
openai: 'gpt-5.5',
|
||||
anthropic: 'claude-sonnet-4-6',
|
||||
gemini: 'gemini-2.5-pro',
|
||||
antigravity: 'claude-sonnet-4-6',
|
||||
grok: 'grok-4.5',
|
||||
kimi: 'kimi-k2.5',
|
||||
zhipu: 'glm-4.7',
|
||||
deepseek: 'deepseek-v4-pro',
|
||||
composite: 'gpt-5.5'
|
||||
}
|
||||
const preferredModel = preferredModels[platform] || ''
|
||||
const model = selectCodexCatalogModel(preferredModel)
|
||||
const labels: Record<GroupPlatform, string> = {
|
||||
anthropic: 'Anthropic',
|
||||
openai: 'OpenAI',
|
||||
gemini: 'Gemini',
|
||||
antigravity: 'Antigravity',
|
||||
grok: 'Grok',
|
||||
kimi: 'Kimi',
|
||||
zhipu: 'Zhipu',
|
||||
deepseek: 'DeepSeek',
|
||||
composite: 'Composite'
|
||||
}
|
||||
const label = labels[platform]
|
||||
const envContent = isWindows
|
||||
? `$env:SUB2API_API_KEY="${apiKey}"`
|
||||
: `export SUB2API_API_KEY="${apiKey}"`
|
||||
|
||||
const configContent = `# Codex CLI -> Sub2API ${label} group
|
||||
model_provider = "sub2api"
|
||||
model = "${model}"
|
||||
review_model = "${model}"
|
||||
disable_response_storage = true
|
||||
model_catalog_json = "${escapeTomlBasicString(codexModelCatalogPath.value)}"
|
||||
|
||||
[model_providers.sub2api]
|
||||
name = "Sub2API ${label}"
|
||||
base_url = "${baseUrl}"
|
||||
env_key = "SUB2API_API_KEY"
|
||||
wire_api = "responses"
|
||||
requires_openai_auth = false
|
||||
supports_websockets = false`
|
||||
|
||||
return [
|
||||
{ path: isWindows ? 'PowerShell' : 'Terminal', content: envContent },
|
||||
{
|
||||
path: joinConfigPath(configDir, 'config.toml', isWindows),
|
||||
content: configContent,
|
||||
hint: t(
|
||||
platform === 'deepseek' || platform === 'composite'
|
||||
? `keys.useKeyModal.${platform}.codexConfigTomlHint`
|
||||
: 'keys.useKeyModal.routedCodex.configTomlHint'
|
||||
)
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
function generateOpenAIWsFiles(baseUrl: string, apiKey: string): FileConfig[] {
|
||||
const isWindows = activeTab.value === 'windows'
|
||||
const configDir = isWindows ? '%userprofile%\\.codex' : '~/.codex'
|
||||
const model = selectCodexCatalogModel('gpt-5.5')
|
||||
const reasoningEffortLine = codexReasoningEffortTomlLine(model)
|
||||
|
||||
// config.toml content with WebSocket v2
|
||||
const configContent = `model_provider = "OpenAI"
|
||||
model = "gpt-5.5"
|
||||
review_model = "gpt-5.5"
|
||||
model_reasoning_effort = "xhigh"
|
||||
disable_response_storage = true
|
||||
model = "${model}"
|
||||
review_model = "${model}"
|
||||
${reasoningEffortLine}disable_response_storage = true
|
||||
model_catalog_json = "${escapeTomlBasicString(codexModelCatalogPath.value)}"
|
||||
network_access = "enabled"
|
||||
windows_wsl_setup_acknowledged = true
|
||||
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
import { describe, expect, it, vi } from 'vitest'
|
||||
import { mount } from '@vue/test-utils'
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest'
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { nextTick } from 'vue'
|
||||
|
||||
const { copyToClipboardMock } = vi.hoisted(() => ({
|
||||
copyToClipboardMock: vi.fn().mockResolvedValue(true)
|
||||
const { copyToClipboardMock, saveAsMock } = vi.hoisted(() => ({
|
||||
copyToClipboardMock: vi.fn().mockResolvedValue(true),
|
||||
saveAsMock: vi.fn()
|
||||
}))
|
||||
|
||||
vi.mock('vue-i18n', () => ({
|
||||
@@ -18,9 +19,26 @@ vi.mock('@/composables/useClipboard', () => ({
|
||||
})
|
||||
}))
|
||||
|
||||
vi.mock('file-saver', () => ({
|
||||
saveAs: saveAsMock
|
||||
}))
|
||||
|
||||
import UseKeyModal from '../UseKeyModal.vue'
|
||||
|
||||
function readBlobAsText(blob: Blob): Promise<string> {
|
||||
return new Promise((resolve, reject) => {
|
||||
const reader = new FileReader()
|
||||
reader.addEventListener('load', () => resolve(String(reader.result || '')))
|
||||
reader.addEventListener('error', () => reject(reader.error))
|
||||
reader.readAsText(blob)
|
||||
})
|
||||
}
|
||||
|
||||
describe('UseKeyModal', () => {
|
||||
afterEach(() => {
|
||||
vi.unstubAllGlobals()
|
||||
saveAsMock.mockClear()
|
||||
})
|
||||
it('renders Grok Build and OpenCode setup for Grok groups', async () => {
|
||||
const wrapper = mount(UseKeyModal, {
|
||||
props: {
|
||||
@@ -299,6 +317,7 @@ describe('UseKeyModal', () => {
|
||||
expect(configToml).not.toContain('supports_websockets')
|
||||
expect(configToml).not.toContain('responses_websockets_v2')
|
||||
expect(configToml).toContain('[features]\ngoals = true')
|
||||
expect(configToml).not.toContain('model_reasoning_effort = "xhigh"')
|
||||
expect(codeBlocks).toContain('{\n "OPENAI_API_KEY": "sk-test"\n}')
|
||||
expect(wrapper.text()).toContain('auth.json')
|
||||
expect(wrapper.find('[data-testid="codex-api-key-restart-notice"]').exists()).toBe(false)
|
||||
@@ -594,4 +613,234 @@ describe('UseKeyModal', () => {
|
||||
expect(fable.options.thinking).toEqual({ type: 'adaptive' })
|
||||
expect(fable.options.thinking).not.toHaveProperty('budgetTokens')
|
||||
})
|
||||
|
||||
// Scenario: API Key users can fetch a routed group catalog and reference it from config.toml.
|
||||
it('offers a downloadable Codex catalog for Composite API keys', async () => {
|
||||
const manifest = {
|
||||
models: [
|
||||
{
|
||||
slug: 'claude-opus-4-8',
|
||||
default_reasoning_level: 'medium',
|
||||
supported_reasoning_levels: [{ effort: 'max', description: 'Maximum reasoning depth' }],
|
||||
input_modalities: ['text'],
|
||||
model_messages: { instructions_template: 'Use the routed model.' }
|
||||
},
|
||||
{
|
||||
slug: 'grok-4.6',
|
||||
default_reasoning_level: 'high',
|
||||
supported_reasoning_levels: [{ effort: 'xhigh', description: 'Extra-high 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 wrapper = mount(UseKeyModal, {
|
||||
props: {
|
||||
show: true,
|
||||
apiKey: 'sk-composite-test',
|
||||
baseUrl: 'https://example.com/v1',
|
||||
platform: 'composite'
|
||||
},
|
||||
global: {
|
||||
stubs: {
|
||||
BaseDialog: {
|
||||
template: '<div><slot /><slot name="footer" /></div>'
|
||||
},
|
||||
Icon: {
|
||||
template: '<span />'
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const codexTab = wrapper.findAll('button').find((button) =>
|
||||
button.text().includes('keys.useKeyModal.cliTabs.codexCli')
|
||||
)
|
||||
expect(codexTab).toBeDefined()
|
||||
await codexTab!.trigger('click')
|
||||
await nextTick()
|
||||
|
||||
const unixConfig = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('[model_providers.sub2api]'))
|
||||
expect(unixConfig).toContain('model_catalog_json = "~/.codex/codex-models.json"')
|
||||
expect(unixConfig).toContain('env_key = "SUB2API_API_KEY"')
|
||||
|
||||
await wrapper.get('[data-testid="codex-model-catalog-fetch"]').trigger('click')
|
||||
await flushPromises()
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledWith(
|
||||
'https://example.com/v1/models?client_version=0.147.0',
|
||||
expect.objectContaining({
|
||||
headers: expect.objectContaining({ Authorization: 'Bearer sk-composite-test' })
|
||||
})
|
||||
)
|
||||
expect(wrapper.get('[data-testid="codex-model-catalog"]').text())
|
||||
.toContain('keys.useKeyModal.codexModelCatalog.download')
|
||||
|
||||
const loadedUnixConfig = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('[model_providers.sub2api]'))
|
||||
expect(loadedUnixConfig).toContain('model = "claude-opus-4-8"')
|
||||
expect(loadedUnixConfig).toContain('review_model = "claude-opus-4-8"')
|
||||
expect(loadedUnixConfig).not.toContain('model = "gpt-5.5"')
|
||||
|
||||
const downloadButton = wrapper.findAll('button').find((button) =>
|
||||
button.text().includes('keys.useKeyModal.codexModelCatalog.download')
|
||||
)
|
||||
expect(downloadButton).toBeDefined()
|
||||
await downloadButton!.trigger('click')
|
||||
expect(saveAsMock).toHaveBeenCalledWith(expect.any(Blob), 'codex-models.json')
|
||||
const downloadedBlob = saveAsMock.mock.calls[0]?.[0] as Blob
|
||||
expect(JSON.parse(await readBlobAsText(downloadedBlob))).toEqual(manifest)
|
||||
|
||||
const windowsTab = wrapper.findAll('button').find((button) => button.text().trim() === 'Windows')
|
||||
expect(windowsTab).toBeDefined()
|
||||
await windowsTab!.trigger('click')
|
||||
await nextTick()
|
||||
|
||||
const windowsConfig = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('[model_providers.sub2api]'))
|
||||
expect(windowsConfig).toContain(
|
||||
'model_catalog_json = "%userprofile%\\\\.codex\\\\codex-models.json"'
|
||||
)
|
||||
})
|
||||
|
||||
it.each(['anthropic', 'gemini', 'antigravity', 'kimi', 'zhipu'] as const)(
|
||||
'offers Codex catalog configuration for the %s routed group',
|
||||
async (platform) => {
|
||||
const wrapper = mount(UseKeyModal, {
|
||||
props: {
|
||||
show: true,
|
||||
apiKey: `sk-${platform}-test`,
|
||||
baseUrl: 'https://example.com/v1',
|
||||
platform
|
||||
},
|
||||
global: {
|
||||
stubs: {
|
||||
BaseDialog: {
|
||||
template: '<div><slot /><slot name="footer" /></div>'
|
||||
},
|
||||
Icon: {
|
||||
template: '<span />'
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const codexTab = wrapper.findAll('button').find((button) =>
|
||||
button.text().includes('keys.useKeyModal.cliTabs.codexCli')
|
||||
)
|
||||
expect(codexTab).toBeDefined()
|
||||
await codexTab!.trigger('click')
|
||||
await nextTick()
|
||||
|
||||
expect(wrapper.find('[data-testid="codex-model-catalog"]').exists()).toBe(true)
|
||||
const config = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('[model_providers.sub2api]'))
|
||||
expect(config).toContain('model_catalog_json = "~/.codex/codex-models.json"')
|
||||
expect(config).toContain('base_url = "https://example.com/v1"')
|
||||
expect(config).toContain('wire_api = "responses"')
|
||||
}
|
||||
)
|
||||
|
||||
// Scenario: the platform-preferred model remains selected when the downloaded catalog contains it.
|
||||
it('keeps the preferred Composite default when it exists in the catalog', async () => {
|
||||
vi.stubGlobal('fetch', vi.fn().mockResolvedValue({
|
||||
ok: true,
|
||||
status: 200,
|
||||
json: async () => ({
|
||||
models: [
|
||||
{ slug: 'claude-opus-4-8' },
|
||||
{ slug: 'gpt-5.5' }
|
||||
]
|
||||
})
|
||||
}))
|
||||
|
||||
const wrapper = mount(UseKeyModal, {
|
||||
props: {
|
||||
show: true,
|
||||
apiKey: 'sk-composite-test',
|
||||
baseUrl: 'https://example.com/v1',
|
||||
platform: 'composite'
|
||||
},
|
||||
global: {
|
||||
stubs: {
|
||||
BaseDialog: {
|
||||
template: '<div><slot /><slot name="footer" /></div>'
|
||||
},
|
||||
Icon: {
|
||||
template: '<span />'
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const codexTab = wrapper.findAll('button').find((button) =>
|
||||
button.text().includes('keys.useKeyModal.cliTabs.codexCli')
|
||||
)
|
||||
expect(codexTab).toBeDefined()
|
||||
await codexTab!.trigger('click')
|
||||
await wrapper.get('[data-testid="codex-model-catalog-fetch"]').trigger('click')
|
||||
await flushPromises()
|
||||
|
||||
const config = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('[model_providers.sub2api]'))
|
||||
expect(config).toContain('model = "gpt-5.5"')
|
||||
expect(config).toContain('review_model = "gpt-5.5"')
|
||||
})
|
||||
|
||||
it('derives OpenAI Codex reasoning effort from the selected catalog descriptor', async () => {
|
||||
vi.stubGlobal('fetch', vi.fn().mockResolvedValue({
|
||||
ok: true,
|
||||
status: 200,
|
||||
json: async () => ({
|
||||
models: [
|
||||
{
|
||||
slug: 'glm-5.3',
|
||||
default_reasoning_level: 'none',
|
||||
supported_reasoning_levels: [{ effort: 'none' }]
|
||||
}
|
||||
]
|
||||
})
|
||||
}))
|
||||
|
||||
const wrapper = mount(UseKeyModal, {
|
||||
props: {
|
||||
show: true,
|
||||
apiKey: 'sk-openai-test',
|
||||
baseUrl: 'https://example.com/v1',
|
||||
platform: 'openai'
|
||||
},
|
||||
global: {
|
||||
stubs: {
|
||||
BaseDialog: {
|
||||
template: '<div><slot /><slot name="footer" /></div>'
|
||||
},
|
||||
Icon: {
|
||||
template: '<span />'
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
await wrapper.get('[data-testid="codex-model-catalog-fetch"]').trigger('click')
|
||||
await flushPromises()
|
||||
|
||||
const configToml = wrapper.findAll('pre code')
|
||||
.map((code) => code.text())
|
||||
.find((content) => content.includes('model_provider = "OpenAI"'))
|
||||
expect(configToml).toContain('model = "glm-5.3"')
|
||||
expect(configToml).not.toContain('model_reasoning_effort')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -743,6 +743,8 @@ export default {
|
||||
syncUpstreamModelsEmpty: 'Upstream returned no models to sync',
|
||||
syncUpstreamModelsFailed: 'Failed to sync upstream models',
|
||||
syncUpstreamModelsError: 'Failed to sync upstream models: {message}',
|
||||
syncUpstreamModelsMetadataIncomplete:
|
||||
'Model IDs were synced, but capability metadata is incomplete and was not updated.',
|
||||
clearAllModels: 'Clear all models',
|
||||
customModelName: 'Custom model name',
|
||||
enterCustomModelName: 'Enter custom model name',
|
||||
|
||||
@@ -188,6 +188,32 @@ export default {
|
||||
codexNoteWindows:
|
||||
'Set $env:SUB2API_API_KEY, save config.toml under %USERPROFILE%\\.codex. Prefer env_key auth; do not commit secrets.',
|
||||
},
|
||||
deepseek: {
|
||||
description: 'Configure Claude Code, Codex, or OpenCode through the current DeepSeek group.',
|
||||
codexDescription: 'Configure Codex with API key authentication through the current DeepSeek group.',
|
||||
codexConfigTomlHint: 'Download the model catalog below, save both files under the Codex config directory, and restart Codex.',
|
||||
codexNote: 'Export SUB2API_API_KEY before starting Codex. The downloaded catalog contains model metadata only, not your API key.',
|
||||
},
|
||||
composite: {
|
||||
description: 'Configure supported clients through the current Composite routing group.',
|
||||
codexDescription: 'Configure Codex with API key authentication and the complete model catalog for this Composite group.',
|
||||
codexConfigTomlHint: 'Download the model catalog below, save both files under the Codex config directory, and restart Codex.',
|
||||
codexNote: 'Export SUB2API_API_KEY before starting Codex. Model requests are routed by the selected catalog slug.',
|
||||
},
|
||||
routedCodex: {
|
||||
description: 'Configure Codex with the complete model catalog for the current routed group.',
|
||||
configTomlHint: 'Download the model catalog below, save both files under the Codex config directory, and restart Codex.',
|
||||
note: 'Export SUB2API_API_KEY before starting Codex. The downloaded catalog contains model metadata only, not your API key.',
|
||||
},
|
||||
codexModelCatalog: {
|
||||
title: 'Codex model catalog',
|
||||
description: 'Fetch with this API key, then save the catalog at the path referenced by config.toml.',
|
||||
fetch: 'Fetch catalog',
|
||||
retry: 'Retry',
|
||||
download: 'Download catalog',
|
||||
modelsCount: '{count} models ready to download',
|
||||
errorDescription: 'The catalog could not be fetched with this API key.',
|
||||
},
|
||||
opencode: {
|
||||
title: 'OpenCode Example',
|
||||
subtitle: 'opencode.json',
|
||||
|
||||
@@ -819,6 +819,7 @@ export default {
|
||||
syncUpstreamModelsEmpty: '上游没有返回可同步的模型',
|
||||
syncUpstreamModelsFailed: '同步上游模型失败',
|
||||
syncUpstreamModelsError: '同步上游模型失败:{message}',
|
||||
syncUpstreamModelsMetadataIncomplete: '模型 ID 已同步,但能力元数据不完整,能力信息未更新。',
|
||||
clearAllModels: '清除所有模型',
|
||||
customModelName: '自定义模型名称',
|
||||
enterCustomModelName: '输入自定义模型名称',
|
||||
|
||||
@@ -192,6 +192,32 @@ export default {
|
||||
codexNoteWindows:
|
||||
'设置 $env:SUB2API_API_KEY,将 config.toml 保存到 %USERPROFILE%\\.codex。优先 env_key,勿提交密钥。'
|
||||
},
|
||||
deepseek: {
|
||||
description: '通过当前 DeepSeek 分组配置 Claude Code、Codex 或 OpenCode。',
|
||||
codexDescription: '使用 API Key 配置 Codex,并通过当前 DeepSeek 分组发送请求。',
|
||||
codexConfigTomlHint: '下载下方模型目录,将两个文件保存到 Codex 配置目录后重启 Codex。',
|
||||
codexNote: '启动 Codex 前先导出 SUB2API_API_KEY。下载的目录只包含模型元数据,不包含 API Key。'
|
||||
},
|
||||
composite: {
|
||||
description: '通过当前 Composite 路由分组配置受支持的客户端。',
|
||||
codexDescription: '使用 API Key 和当前 Composite 分组的完整模型目录配置 Codex。',
|
||||
codexConfigTomlHint: '下载下方模型目录,将两个文件保存到 Codex 配置目录后重启 Codex。',
|
||||
codexNote: '启动 Codex 前先导出 SUB2API_API_KEY;分组会根据目录中选中的模型路由请求。'
|
||||
},
|
||||
routedCodex: {
|
||||
description: '使用当前路由分组的完整模型目录配置 Codex。',
|
||||
configTomlHint: '下载下方模型目录,将两个文件保存到 Codex 配置目录后重启 Codex。',
|
||||
note: '启动 Codex 前先导出 SUB2API_API_KEY。下载的目录只包含模型元数据,不包含 API Key。'
|
||||
},
|
||||
codexModelCatalog: {
|
||||
title: 'Codex 模型目录',
|
||||
description: '使用当前 API Key 获取目录,并保存到 config.toml 引用的路径。',
|
||||
fetch: '获取目录',
|
||||
retry: '重试',
|
||||
download: '下载目录',
|
||||
modelsCount: '已获取 {count} 个模型',
|
||||
errorDescription: '无法使用当前 API Key 获取模型目录。'
|
||||
},
|
||||
opencode: {
|
||||
title: 'OpenCode 配置示例',
|
||||
subtitle: 'opencode.json',
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import {
|
||||
findCodexCatalogModel,
|
||||
formatCodexReasoningEffortTomlLine,
|
||||
parseCodexCatalogModels,
|
||||
selectCodexConfigReasoningEffort
|
||||
} from '@/utils/codexCatalogConfig'
|
||||
|
||||
describe('codexCatalogConfig', () => {
|
||||
it('parses catalog slugs and finds a model by id', () => {
|
||||
const content = JSON.stringify({
|
||||
models: [
|
||||
{ slug: 'glm-5.3', default_reasoning_level: 'none', supported_reasoning_levels: [{ effort: 'none' }] },
|
||||
{ slug: ' ', supported_reasoning_levels: [] }
|
||||
]
|
||||
})
|
||||
expect(parseCodexCatalogModels(content).map((model) => model.slug)).toEqual(['glm-5.3'])
|
||||
expect(findCodexCatalogModel(content, 'glm-5.3')?.slug).toBe('glm-5.3')
|
||||
expect(findCodexCatalogModel(content, 'missing')).toBeUndefined()
|
||||
})
|
||||
|
||||
it('omits effort when the descriptor only advertises none', () => {
|
||||
expect(selectCodexConfigReasoningEffort({
|
||||
slug: 'glm-5.3',
|
||||
default_reasoning_level: 'none',
|
||||
supported_reasoning_levels: [{ effort: 'none' }]
|
||||
})).toBeNull()
|
||||
expect(formatCodexReasoningEffortTomlLine(null)).toBe('')
|
||||
})
|
||||
|
||||
it('does not emit an effort absent from supported_reasoning_levels', () => {
|
||||
expect(selectCodexConfigReasoningEffort({
|
||||
slug: 'glm-5.3',
|
||||
default_reasoning_level: 'xhigh',
|
||||
supported_reasoning_levels: [{ effort: 'none' }]
|
||||
})).toBeNull()
|
||||
})
|
||||
|
||||
it('uses the catalog default when it is a supported non-none effort', () => {
|
||||
expect(selectCodexConfigReasoningEffort({
|
||||
slug: 'gpt-5.5',
|
||||
default_reasoning_level: 'medium',
|
||||
supported_reasoning_levels: [
|
||||
{ effort: 'low' },
|
||||
{ effort: 'medium' },
|
||||
{ effort: 'high' },
|
||||
{ effort: 'xhigh' }
|
||||
]
|
||||
})).toBe('medium')
|
||||
expect(formatCodexReasoningEffortTomlLine('medium')).toBe('model_reasoning_effort = "medium"\n')
|
||||
})
|
||||
|
||||
it('falls back to the first usable supported effort when default is missing', () => {
|
||||
expect(selectCodexConfigReasoningEffort({
|
||||
slug: 'custom',
|
||||
supported_reasoning_levels: [{ effort: 'none' }, { effort: 'high' }]
|
||||
})).toBe('high')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,63 @@
|
||||
export interface CodexCatalogReasoningLevel {
|
||||
effort?: unknown
|
||||
}
|
||||
|
||||
export interface CodexCatalogModel {
|
||||
slug: string
|
||||
default_reasoning_level?: unknown
|
||||
supported_reasoning_levels?: CodexCatalogReasoningLevel[]
|
||||
}
|
||||
|
||||
function trimEffort(value: unknown): string {
|
||||
if (typeof value !== 'string') return ''
|
||||
return value.trim()
|
||||
}
|
||||
|
||||
export function parseCodexCatalogModels(content: string | null | undefined): CodexCatalogModel[] {
|
||||
if (!content) return []
|
||||
try {
|
||||
const payload: unknown = JSON.parse(content)
|
||||
if (typeof payload !== 'object' || payload === null || !('models' in payload)) return []
|
||||
const models = (payload as { models?: unknown }).models
|
||||
if (!Array.isArray(models)) return []
|
||||
return models.flatMap((model) => {
|
||||
if (typeof model !== 'object' || model === null || !('slug' in model)) return []
|
||||
const slug = trimEffort((model as { slug?: unknown }).slug)
|
||||
if (!slug) return []
|
||||
return [{ ...(model as CodexCatalogModel), slug }]
|
||||
})
|
||||
} catch {
|
||||
return []
|
||||
}
|
||||
}
|
||||
|
||||
export function findCodexCatalogModel(
|
||||
content: string | null | undefined,
|
||||
slug: string
|
||||
): CodexCatalogModel | undefined {
|
||||
const wanted = slug.trim()
|
||||
if (!wanted) return undefined
|
||||
return parseCodexCatalogModels(content).find((model) => model.slug === wanted)
|
||||
}
|
||||
|
||||
export function selectCodexConfigReasoningEffort(
|
||||
model: CodexCatalogModel | undefined
|
||||
): string | null {
|
||||
if (!model) return null
|
||||
const efforts = (model.supported_reasoning_levels ?? []).flatMap((level) => {
|
||||
const effort = trimEffort(level?.effort)
|
||||
return effort ? [effort] : []
|
||||
})
|
||||
if (efforts.length === 0) return null
|
||||
|
||||
const defaultLevel = trimEffort(model.default_reasoning_level)
|
||||
if (defaultLevel && efforts.includes(defaultLevel)) {
|
||||
return defaultLevel === 'none' ? null : defaultLevel
|
||||
}
|
||||
return efforts.find((effort) => effort !== 'none') ?? null
|
||||
}
|
||||
|
||||
export function formatCodexReasoningEffortTomlLine(effort: string | null): string {
|
||||
if (!effort) return ''
|
||||
return `model_reasoning_effort = "${effort}"\n`
|
||||
}
|
||||
Reference in New Issue
Block a user