mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 14:38:27 +08:00
Merge branch 'main' into fix/openai-raw-stream-truncation
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.181
|
||||
0.1.183
|
||||
|
||||
@@ -105,11 +105,11 @@ var DefaultAntigravityModelMapping = map[string]string{
|
||||
"claude-opus-4-6": "claude-opus-4-6-thinking", // 简称映射
|
||||
"claude-opus-4-5-thinking": "claude-opus-4-6-thinking", // 迁移旧模型
|
||||
"claude-sonnet-4-6": "claude-sonnet-4-6",
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5-thinking": "claude-sonnet-4-5-thinking",
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-5", // 显式 canonical 选择透传
|
||||
"claude-sonnet-4-5-thinking": "claude-sonnet-4-6", // 迁移旧兼容别名
|
||||
// Claude 详细版本 ID 映射
|
||||
"claude-opus-4-5-20251101": "claude-opus-4-6-thinking", // 迁移旧模型
|
||||
"claude-sonnet-4-5-20250929": "claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5-20250929": "claude-sonnet-4-6", // 迁移旧模型
|
||||
// Claude Haiku → Sonnet(无 Haiku 支持)
|
||||
"claude-haiku-4-5": "claude-sonnet-4-6",
|
||||
"claude-haiku-4-5-20251001": "claude-sonnet-4-6",
|
||||
|
||||
@@ -43,6 +43,21 @@ func TestDefaultAntigravityModelMapping_ContainsNewClaudeModels(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAntigravityModelMapping_PreservesExplicitSonnet45AndMigratesLegacyAliases(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := map[string]string{
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5-thinking": "claude-sonnet-4-6",
|
||||
"claude-sonnet-4-5-20250929": "claude-sonnet-4-6",
|
||||
}
|
||||
for model, want := range cases {
|
||||
if got := DefaultAntigravityModelMapping[model]; got != want {
|
||||
t.Fatalf("expected model %q to map to %q, got %q", model, want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -31,6 +31,7 @@ func TestOpenAICompatibleTextTargetAllowsCompositeProviders(t *testing.T) {
|
||||
}{
|
||||
{model: "grok-4.3", platform: service.PlatformGrok},
|
||||
{model: "kimi-k2-thinking", platform: service.PlatformKimi},
|
||||
{model: "k3", platform: service.PlatformKimi},
|
||||
{model: "glm-5.2", platform: service.PlatformZhipu},
|
||||
{model: "deepseek-v3.2", platform: service.PlatformDeepseek},
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -1234,11 +1307,15 @@ func writeGrokModelsList(c *gin.Context, modelIDs []string) {
|
||||
if grokModelSupportsConfigurableReasoning(modelID) {
|
||||
item.SupportsReasoningEffort = true
|
||||
item.ReasoningEffort = "high"
|
||||
item.ReasoningEfforts = []grokReasoningEffortOption{
|
||||
efforts := []grokReasoningEffortOption{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
}
|
||||
if service.GrokSupportsXHighReasoningEffort(modelID) {
|
||||
efforts = append(efforts, grokReasoningEffortOption{Value: "xhigh", Label: "xHigh"})
|
||||
}
|
||||
item.ReasoningEfforts = efforts
|
||||
}
|
||||
models = append(models, item)
|
||||
}
|
||||
@@ -1340,9 +1417,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 +1455,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)
|
||||
|
||||
@@ -103,9 +364,38 @@ func TestGatewayModels_GeminiGroupFallsBackToGeminiModels(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T) {
|
||||
assertGrokGatewayReasoningEfforts(t, 4409, "grok-4.5", []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
})
|
||||
}
|
||||
|
||||
func TestGatewayModels_Grok46AdvertisesXHighReasoningEffortForGrokBuild(t *testing.T) {
|
||||
xhighEfforts := []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
{Value: "xhigh", Label: "xHigh"},
|
||||
}
|
||||
tests := []struct {
|
||||
groupID int64
|
||||
model string
|
||||
}{
|
||||
{groupID: 4410, model: "grok-4.6"},
|
||||
{groupID: 4411, model: "grok-4.6-latest"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.model, func(t *testing.T) {
|
||||
assertGrokGatewayReasoningEfforts(t, tt.groupID, tt.model, xhighEfforts)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertGrokGatewayReasoningEfforts(t *testing.T, groupID int64, modelID string, want []gatewayReasoningEffortOptionForTest) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
groupID := int64(4409)
|
||||
h := newGatewayModelsHandlerForTest(
|
||||
&gatewayModelsAccountRepoStub{
|
||||
byGroup: map[int64][]service.Account{
|
||||
@@ -114,7 +404,7 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
|
||||
ID: 1,
|
||||
Platform: service.PlatformGrok,
|
||||
Credentials: map[string]any{
|
||||
"model_mapping": map[string]any{"grok-4.5": "grok-4.5"},
|
||||
"model_mapping": map[string]any{modelID: modelID},
|
||||
},
|
||||
},
|
||||
},
|
||||
@@ -136,14 +426,10 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
require.Len(t, got.Data, 1)
|
||||
model := got.Data[0]
|
||||
require.Equal(t, "grok-4.5", model.ID)
|
||||
require.Equal(t, modelID, model.ID)
|
||||
require.True(t, model.SupportsReasoningEffort)
|
||||
require.Equal(t, "high", model.ReasoningEffort)
|
||||
require.Equal(t, []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
}, model.ReasoningEfforts)
|
||||
require.Equal(t, want, model.ReasoningEfforts)
|
||||
}
|
||||
|
||||
func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) {
|
||||
@@ -193,6 +479,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 +799,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
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -19,13 +20,16 @@ type responsesFailedError struct {
|
||||
|
||||
// responsesFailedBody 对齐 apicompat.makeResponsesCompletedEvent 输出的 response 子对象字段集。
|
||||
// Output 用空 slice(不是 nil)确保 marshal 为 `[]` 而非 `null`。
|
||||
// CreatedAt 不带 omitempty:严格客户端把它当必填字段,缺失会以
|
||||
// `missing field 'created_at'` 反序列化失败——那正是本文件要避免的"客户端读不懂终止事件"。
|
||||
type responsesFailedBody struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Output []any `json:"output"`
|
||||
Error responsesFailedError `json:"error"`
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Output []any `json:"output"`
|
||||
Error responsesFailedError `json:"error"`
|
||||
}
|
||||
|
||||
// responsesFailedEvent 是写入 SSE data 行的顶层结构。
|
||||
@@ -61,11 +65,12 @@ func writeResponsesFailedSSE(c *gin.Context, errType, message string) bool {
|
||||
payload, err := json.Marshal(responsesFailedEvent{
|
||||
Type: "response.failed",
|
||||
Response: responsesFailedBody{
|
||||
ID: synthesizeResponseID(c),
|
||||
Object: "response",
|
||||
Model: requestModel(c),
|
||||
Status: "failed",
|
||||
Output: []any{},
|
||||
ID: synthesizeResponseID(c),
|
||||
Object: "response",
|
||||
CreatedAt: time.Now().Unix(),
|
||||
Model: requestModel(c),
|
||||
Status: "failed",
|
||||
Output: []any{},
|
||||
Error: responsesFailedError{
|
||||
Code: mapResponsesErrorCode(errType),
|
||||
Message: message,
|
||||
|
||||
@@ -227,6 +227,22 @@ func TestOpenAIHandleStreamingAwareError_BareResponsesRouteEmitsResponseFailed(t
|
||||
}
|
||||
|
||||
// Synthesized response.failed id falls back to uuid when no request_id is present.
|
||||
// issue #5601:严格的 Responses 客户端把 created_at 当必填字段,缺失即
|
||||
// `missing field 'created_at'`。合成的终止事件若解析不了,本文件存在的意义
|
||||
// (给客户端一个可识别的终止事件而不是盲重连)就落空了。
|
||||
func TestOpenAIHandleStreamingAwareError_ResponsesStreamingCarriesCreatedAt(t *testing.T) {
|
||||
c, w := newGinContextForEndpoint(t, EndpointResponses)
|
||||
h := &OpenAIGatewayHandler{}
|
||||
h.handleStreamingAwareError(c, http.StatusBadGateway, "upstream_error", "boom", true)
|
||||
|
||||
resp, _ := parseResponsesFailedSSE(t, w.Body.String())
|
||||
raw, ok := resp["created_at"]
|
||||
assert.True(t, ok, "response.failed 必须带 created_at")
|
||||
createdAt, ok := raw.(float64)
|
||||
assert.True(t, ok, "created_at 必须是数字,得到 %T", raw)
|
||||
assert.Greater(t, int64(createdAt), int64(0), "created_at 必须是有效的 unix 时间戳")
|
||||
}
|
||||
|
||||
func TestSynthesizeResponseID_FallbackUUID(t *testing.T) {
|
||||
c, _ := newGinContextForEndpoint(t, EndpointResponses)
|
||||
id := synthesizeResponseID(c)
|
||||
|
||||
@@ -21,10 +21,13 @@ func AnthropicToResponsesResponse(resp *AnthropicResponse) *ResponsesResponse {
|
||||
id = generateResponsesID()
|
||||
}
|
||||
|
||||
// Anthropic responses carry no creation timestamp, so stamp now — the same
|
||||
// synthesize-what-the-client-requires rule the generated id above follows.
|
||||
out := &ResponsesResponse{
|
||||
ID: id,
|
||||
Object: "response",
|
||||
Model: resp.Model,
|
||||
ID: id,
|
||||
Object: "response",
|
||||
CreatedAt: time.Now().Unix(),
|
||||
Model: resp.Model,
|
||||
}
|
||||
|
||||
var outputs []ResponsesOutput
|
||||
@@ -551,11 +554,12 @@ func makeResponsesCreatedEvent(state *AnthropicEventToResponsesState) ResponsesS
|
||||
Type: "response.created",
|
||||
SequenceNumber: seq,
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
Output: []ResponsesOutput{},
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
Output: []ResponsesOutput{},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -602,6 +606,7 @@ func makeResponsesCompletedEvent(
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: status,
|
||||
Output: outputs,
|
||||
|
||||
@@ -249,8 +249,8 @@ func anthropicUserToChatMessages(raw json.RawMessage) ([]ChatMessage, error) {
|
||||
|
||||
// anthropicAssistantToChatMessages handles an Anthropic assistant message.
|
||||
// Text content → assistant message content; tool_use blocks → tool_calls on the
|
||||
// same assistant message; thinking blocks are dropped (Chat Completions has no
|
||||
// inbound thinking field, matching anthropicAssistantToResponses).
|
||||
// same assistant message; thinking blocks → reasoning_content, but only on a
|
||||
// message that carries tool calls (see anthropicThinkingToReasoningContent).
|
||||
func anthropicAssistantToChatMessages(raw json.RawMessage) ([]ChatMessage, error) {
|
||||
// Plain string → single assistant message.
|
||||
var s string
|
||||
@@ -289,9 +289,40 @@ func anthropicAssistantToChatMessages(raw json.RawMessage) ([]ChatMessage, error
|
||||
})
|
||||
}
|
||||
|
||||
msg.ReasoningContent = anthropicThinkingToReasoningContent(blocks, len(msg.ToolCalls) > 0)
|
||||
|
||||
return []ChatMessage{msg}, nil
|
||||
}
|
||||
|
||||
// anthropicThinkingToReasoningContent folds thinking blocks back into the
|
||||
// Chat Completions reasoning_content field.
|
||||
//
|
||||
// chatMessageToAnthropicBlocks emits the upstream's reasoning_content as a
|
||||
// thinking block on the way out, so a multi-turn client echoes it back on the
|
||||
// next request; dropping it here made the bridge lose exactly what it had just
|
||||
// produced. DeepSeek's thinking mode requires the reasoning_content that
|
||||
// produced a tool call to be replayed on that assistant message and answers
|
||||
// 400 otherwise, which is why buildChatMessagesFromItems already carries
|
||||
// pendingReasoning onto assistant tool-call messages in the Responses→Chat
|
||||
// bridge. hasToolCalls keeps the scope identical to that sibling: reasoning
|
||||
// rides along with tool calls only, never on a plain assistant text turn.
|
||||
//
|
||||
// redacted_thinking blocks and signature-only placeholders carry no plaintext
|
||||
// and contribute nothing. Multiple blocks join with "\n", matching
|
||||
// extractResponsesReasoningText.
|
||||
func anthropicThinkingToReasoningContent(blocks []AnthropicContentBlock, hasToolCalls bool) string {
|
||||
if !hasToolCalls {
|
||||
return ""
|
||||
}
|
||||
var parts []string
|
||||
for _, b := range blocks {
|
||||
if b.Type == "thinking" && b.Thinking != "" {
|
||||
parts = append(parts, b.Thinking)
|
||||
}
|
||||
}
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
// anthropicToolsToChatTools maps Anthropic tool definitions to Chat Completions
|
||||
// function tools. Server-side tools (web_search_*) are dropped — they have no
|
||||
// Chat Completions equivalent.
|
||||
|
||||
@@ -143,8 +143,11 @@ func TestAnthropicToChatCompletionsRequest_ThinkingDropped(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 1)
|
||||
// Only text survives; thinking is dropped
|
||||
// Only text survives. Thinking is dropped because this turn carries no tool
|
||||
// calls — reasoning rides along with tool calls only, matching the
|
||||
// Responses→Chat bridge (see anthropicThinkingToReasoningContent).
|
||||
require.Equal(t, `"answer"`, string(out.Messages[0].Content))
|
||||
require.Empty(t, out.Messages[0].ReasoningContent)
|
||||
}
|
||||
|
||||
func TestAnthropicToChatCompletionsRequest_ToolChoiceAuto(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,229 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5528:/v1/messages 客户端(Claude Code 等)打到只会 Chat Completions 的
|
||||
// OpenAI 兼容上游时,历史 assistant 消息里的 thinking 块被整块丢弃。DeepSeek 的
|
||||
// thinking mode 要求产生工具调用的 reasoning_content 随该 assistant 消息回传,
|
||||
// 于是「单轮正常、一进多轮工具对话必现 400」。
|
||||
|
||||
func anthropicAssistantMsg(t *testing.T, blocks string) *AnthropicRequest {
|
||||
t.Helper()
|
||||
return &AnthropicRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
MaxTokens: 256,
|
||||
Messages: []AnthropicMessage{
|
||||
{Role: "user", Content: json.RawMessage(`"what's the weather?"`)},
|
||||
{Role: "assistant", Content: json.RawMessage(blocks)},
|
||||
{Role: "user", Content: json.RawMessage(`[{"type":"tool_result","tool_use_id":"toolu_1","content":"sunny"}]`)},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
const anthropicThinkingToolTurn = `[
|
||||
{"type":"thinking","thinking":"user wants weather, call the tool"},
|
||||
{"type":"text","text":"checking"},
|
||||
{"type":"tool_use","id":"toolu_1","name":"get_weather","input":{"city":"SF"}}
|
||||
]`
|
||||
|
||||
func TestAnthropicToChatCompletionsRequest_ThinkingBecomesReasoningContentOnToolTurn(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, anthropicThinkingToolTurn))
|
||||
require.NoError(t, err)
|
||||
|
||||
var assistant *ChatMessage
|
||||
for i := range out.Messages {
|
||||
if out.Messages[i].Role == "assistant" {
|
||||
assistant = &out.Messages[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, assistant, "assistant message must survive the bridge")
|
||||
require.Equal(t, "user wants weather, call the tool", assistant.ReasoningContent,
|
||||
"产生工具调用的 thinking 必须作为 reasoning_content 回传,否则 DeepSeek 400")
|
||||
require.Len(t, assistant.ToolCalls, 1)
|
||||
require.Equal(t, `"checking"`, string(assistant.Content), "text/tool_use 处理保持不变")
|
||||
}
|
||||
|
||||
// 上游线格式才是上游看到的东西:字段没序列化出去,等于没修。
|
||||
func TestAnthropicToChatCompletionsRequest_ReasoningContentSerializesOnWire(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, anthropicThinkingToolTurn))
|
||||
require.NoError(t, err)
|
||||
|
||||
payload, err := json.Marshal(out)
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, string(payload), `"reasoning_content":"user wants weather, call the tool"`)
|
||||
}
|
||||
|
||||
// 闭环不变式:thinking 块本来就是本桥出站时用上游 reasoning_content 生成的
|
||||
// (chatMessageToAnthropicBlocks),客户端只是原样回传。出站造、入站丢 = 自己丢自己的东西。
|
||||
func TestAnthropicChatBridge_ReasoningSurvivesOutboundInboundRoundTrip(t *testing.T) {
|
||||
upstream := ChatMessage{
|
||||
Role: "assistant",
|
||||
ReasoningContent: "step 1: need the weather tool",
|
||||
Content: json.RawMessage(`"checking"`),
|
||||
ToolCalls: []ChatToolCall{{
|
||||
ID: "call_1",
|
||||
Type: "function",
|
||||
Function: ChatFunctionCall{Name: "get_weather", Arguments: `{"city":"SF"}`},
|
||||
}},
|
||||
}
|
||||
|
||||
// 出站:Chat 响应 → Anthropic content blocks
|
||||
blocks := chatMessageToAnthropicBlocks(upstream)
|
||||
require.Equal(t, "thinking", blocks[0].Type)
|
||||
require.Equal(t, upstream.ReasoningContent, blocks[0].Thinking)
|
||||
|
||||
// 客户端下一轮把同一组 blocks 原样回传
|
||||
raw, err := json.Marshal(blocks)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 入站:Anthropic content blocks → Chat 请求
|
||||
back, err := anthropicAssistantToChatMessages(raw)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, back, 1)
|
||||
require.Equal(t, upstream.ReasoningContent, back[0].ReasoningContent,
|
||||
"出站生成的 thinking 必须能原样还原回 reasoning_content")
|
||||
require.Len(t, back[0].ToolCalls, 1)
|
||||
}
|
||||
|
||||
// 兄弟不变式:Responses→Chat 桥(buildChatMessagesFromItems 的 pendingReasoning)
|
||||
// 早就把 reasoning 挂到带 tool_calls 的 assistant 消息上了。等价历史下两条桥必须一致。
|
||||
func TestAnthropicChatBridge_MatchesResponsesChatBridgeReasoningPlacement(t *testing.T) {
|
||||
responsesReq := &ResponsesRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"what's the weather?"}]},
|
||||
{"type":"reasoning","summary":[{"type":"summary_text","text":"call the tool"}]},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_weather","arguments":"{\"city\":\"SF\"}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"sunny"}
|
||||
]`),
|
||||
}
|
||||
viaResponses, err := ResponsesToChatCompletionsRequest(responsesReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
viaAnthropic, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, `[
|
||||
{"type":"thinking","thinking":"call the tool"},
|
||||
{"type":"tool_use","id":"toolu_1","name":"get_weather","input":{"city":"SF"}}
|
||||
]`))
|
||||
require.NoError(t, err)
|
||||
|
||||
reasoningOnToolCallMessage := func(msgs []ChatMessage) string {
|
||||
for _, m := range msgs {
|
||||
if m.Role == "assistant" && len(m.ToolCalls) > 0 {
|
||||
return m.ReasoningContent
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
require.Equal(t, "call the tool", reasoningOnToolCallMessage(viaResponses.Messages),
|
||||
"前置条件:兄弟桥本来就带 reasoning_content")
|
||||
require.Equal(t, reasoningOnToolCallMessage(viaResponses.Messages),
|
||||
reasoningOnToolCallMessage(viaAnthropic.Messages),
|
||||
"两条桥对等价历史必须产出同样的 reasoning_content 位置")
|
||||
}
|
||||
|
||||
// 作用域守卫:不带工具调用的纯文本轮次维持现状(与兄弟桥一致 —— reasoning 只随
|
||||
// 工具调用回传),避免把 reasoning_content 撒到不需要它的上游请求上。
|
||||
func TestAnthropicToChatCompletionsRequest_ThinkingWithoutToolCallsStaysDropped(t *testing.T) {
|
||||
req := &AnthropicRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
MaxTokens: 100,
|
||||
Messages: []AnthropicMessage{
|
||||
{Role: "assistant", Content: json.RawMessage(
|
||||
`[{"type":"thinking","thinking":"secret thoughts"},{"type":"text","text":"answer"}]`)},
|
||||
},
|
||||
}
|
||||
|
||||
out, err := AnthropicToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 1)
|
||||
require.Empty(t, out.Messages[0].ReasoningContent)
|
||||
require.Equal(t, `"answer"`, string(out.Messages[0].Content))
|
||||
|
||||
payload, err := json.Marshal(out)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(payload), "reasoning_content")
|
||||
}
|
||||
|
||||
func TestAnthropicThinkingToReasoningContent(t *testing.T) {
|
||||
blocksOf := func(t *testing.T, raw string) []AnthropicContentBlock {
|
||||
t.Helper()
|
||||
var blocks []AnthropicContentBlock
|
||||
require.NoError(t, json.Unmarshal([]byte(raw), &blocks))
|
||||
return blocks
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
raw string
|
||||
hasToolCalls bool
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "single_thinking_block",
|
||||
raw: `[{"type":"thinking","thinking":"a"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "a",
|
||||
},
|
||||
{
|
||||
// 多个 thinking 块用 "\n" 连接,与 extractResponsesReasoningText 一致。
|
||||
name: "multiple_blocks_join_with_newline",
|
||||
raw: `[{"type":"thinking","thinking":"a"},{"type":"text","text":"x"},{"type":"thinking","thinking":"b"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "a\nb",
|
||||
},
|
||||
{
|
||||
// redacted_thinking 没有明文可回传。
|
||||
name: "redacted_thinking_has_no_plaintext",
|
||||
raw: `[{"type":"redacted_thinking","signature":"abc"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
// 只带 signature 的 thinking 占位块(xAI/Codex 密文回放形态)同样无明文。
|
||||
name: "signature_only_thinking",
|
||||
raw: `[{"type":"thinking","thinking":"","signature":"gAAAAxxx"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "no_tool_calls_returns_empty",
|
||||
raw: `[{"type":"thinking","thinking":"a"}]`,
|
||||
hasToolCalls: false,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "no_thinking_blocks",
|
||||
raw: `[{"type":"text","text":"x"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "empty_blocks",
|
||||
raw: `[]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
require.Equal(t, tc.want,
|
||||
anthropicThinkingToReasoningContent(blocksOf(t, tc.raw), tc.hasToolCalls))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 纯字符串形态的 assistant content 没有 blocks 可读,走早返回分支,不得 panic。
|
||||
func TestAnthropicAssistantToChatMessages_PlainStringContentUnaffected(t *testing.T) {
|
||||
msgs, err := anthropicAssistantToChatMessages(json.RawMessage(`"just text"`))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, msgs, 1)
|
||||
require.Empty(t, msgs[0].ReasoningContent)
|
||||
require.Equal(t, `"just text"`, string(msgs[0].Content))
|
||||
}
|
||||
@@ -1211,9 +1211,20 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str
|
||||
id = generateResponsesID()
|
||||
}
|
||||
|
||||
// Carry the upstream's own creation timestamp when it sent one; otherwise
|
||||
// stamp now, same fallback shape as the generated id above.
|
||||
createdAt := int64(0)
|
||||
if resp != nil {
|
||||
createdAt = resp.Created
|
||||
}
|
||||
if createdAt <= 0 {
|
||||
createdAt = time.Now().Unix()
|
||||
}
|
||||
|
||||
out := &ResponsesResponse{
|
||||
ID: id,
|
||||
Object: "response",
|
||||
CreatedAt: createdAt,
|
||||
Model: model,
|
||||
Status: "completed",
|
||||
ServiceTier: chatServiceTier(resp),
|
||||
@@ -1711,6 +1722,7 @@ func FinalizeChatCompletionsResponsesStream(state *ChatCompletionsToResponsesStr
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: status,
|
||||
ServiceTier: state.ServiceTier,
|
||||
@@ -1731,6 +1743,7 @@ func ensureChatToResponsesCreated(state *ChatCompletionsToResponsesStreamState)
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
ServiceTier: state.ServiceTier,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -299,6 +358,9 @@ func normalizeClientToolOutput(item map[string]any) {
|
||||
if _, ok := output.(string); ok {
|
||||
return
|
||||
}
|
||||
if isResponsesToolOutputContent(output) {
|
||||
return
|
||||
}
|
||||
if output == nil {
|
||||
item["output"] = ""
|
||||
return
|
||||
@@ -311,6 +373,25 @@ func normalizeClientToolOutput(item map[string]any) {
|
||||
item["output"] = string(encoded)
|
||||
}
|
||||
|
||||
func isResponsesToolOutputContent(output any) bool {
|
||||
parts, ok := output.([]any)
|
||||
if !ok || len(parts) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, part := range parts {
|
||||
typed, ok := part.(map[string]any)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
switch stringValue(typed["type"]) {
|
||||
case "input_text", "input_image", "input_file":
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// normalizeToolSearchOutput converts both tool_search output wire shapes into
|
||||
// the string output required by function_call_output. Older clients send an
|
||||
// output field directly; newer Codex clients return discovered definitions in
|
||||
@@ -432,12 +513,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 +547,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 +594,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 +614,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 +637,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 +776,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 +805,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 +866,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])
|
||||
@@ -460,7 +460,55 @@ func TestAdaptResponsesClientToolsWithInheritedMapping_LowersFollowupHistoryWith
|
||||
output := requireResponsesClientToolValue[map[string]any](t, items[1])
|
||||
require.Equal(t, "function_call_output", output["type"])
|
||||
require.NotContains(t, output, "id")
|
||||
require.JSONEq(t, `[{"text":"ok","type":"input_text"}]`, requireResponsesClientToolValue[string](t, output["output"]))
|
||||
require.Equal(t, []any{map[string]any{"type": "input_text", "text": "ok"}}, output["output"])
|
||||
}
|
||||
|
||||
func TestAdaptResponsesClientTools_NormalizesCustomToolOutput(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output any
|
||||
wantOutput any
|
||||
}{
|
||||
{
|
||||
name: "supported content parts remain an array",
|
||||
output: []any{
|
||||
map[string]any{"type": "input_text", "text": "ok"},
|
||||
map[string]any{"type": "input_image", "image_url": "https://example.com/image.png"},
|
||||
map[string]any{"type": "input_file", "file_id": "file_123"},
|
||||
},
|
||||
wantOutput: []any{
|
||||
map[string]any{"type": "input_text", "text": "ok"},
|
||||
map[string]any{"type": "input_image", "image_url": "https://example.com/image.png"},
|
||||
map[string]any{"type": "input_file", "file_id": "file_123"},
|
||||
},
|
||||
},
|
||||
{name: "ordinary object is stringified", output: map[string]any{"ok": true}, wantOutput: `{"ok":true}`},
|
||||
{name: "arbitrary array is stringified", output: []any{"ok"}, wantOutput: `["ok"]`},
|
||||
{name: "empty array is stringified", output: []any{}, wantOutput: `[]`},
|
||||
{name: "mixed array is stringified", output: []any{map[string]any{"type": "input_text", "text": "ok"}, "bad"}, wantOutput: `[{"text":"ok","type":"input_text"},"bad"]`},
|
||||
{name: "unknown content type is stringified", output: []any{map[string]any{"type": "output_text", "text": "bad"}}, wantOutput: `[{"text":"bad","type":"output_text"}]`},
|
||||
{name: "whitespace-padded content type is stringified", output: []any{map[string]any{"type": " input_text ", "text": "bad"}}, wantOutput: `[{"text":"bad","type":" input_text "}]`},
|
||||
{name: "missing content type is stringified", output: []any{map[string]any{"text": "bad"}}, wantOutput: `[{"text":"bad"}]`},
|
||||
{name: "non-string content type is stringified", output: []any{map[string]any{"type": 1}}, wantOutput: `[{"type":1}]`},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
req := map[string]any{
|
||||
"tools": []any{map[string]any{"type": "custom", "name": "exec"}},
|
||||
"input": []any{map[string]any{
|
||||
"type": "custom_tool_call_output", "call_id": "call_1", "output": tc.output,
|
||||
}},
|
||||
}
|
||||
|
||||
_, changed, err := AdaptResponsesClientTools(req)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
item := requireResponsesClientToolValue[map[string]any](t, requireResponsesClientToolValue[[]any](t, req["input"])[0])
|
||||
require.Equal(t, "function_call_output", item["type"])
|
||||
require.Equal(t, tc.wantOutput, item["output"])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdaptResponsesClientToolsWithInheritedMapping_PromotesOmittedToolsDiscoveryIntoEffectiveDeclarations(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5601:严格的 Responses 客户端(Rust serde 系,如 Codex / Grok CLI)把
|
||||
// created_at 声明为必填字段,缺失即 `missing field 'created_at'` 反序列化失败。
|
||||
// 网关合成的 Responses 对象(Chat→Responses、Anthropic→Responses 两座桥)此前从不
|
||||
// 写这个字段——尽管两个流式 state 早就采集好了 Created 时间戳,只是没有出口。
|
||||
// 原生 Responses 透传走 gjson/sjson 字节级改写,不受影响。
|
||||
|
||||
// responseObjectOf 取出事件里的 response 子对象(按线格式,而不是按 Go 结构体)。
|
||||
func responseObjectOf(t *testing.T, evt ResponsesStreamEvent) map[string]any {
|
||||
t.Helper()
|
||||
m := marshalEvent(t, evt)
|
||||
resp, ok := m["response"].(map[string]any)
|
||||
require.True(t, ok, "event must carry a response object: %v", m)
|
||||
return resp
|
||||
}
|
||||
|
||||
func requireCreatedAt(t *testing.T, resp map[string]any) int64 {
|
||||
t.Helper()
|
||||
raw, ok := resp["created_at"]
|
||||
require.True(t, ok, "response 对象必须带 created_at,否则严格客户端直接反序列化失败")
|
||||
value, ok := raw.(float64)
|
||||
require.True(t, ok, "created_at 必须是数字,得到 %T", raw)
|
||||
require.Greater(t, int64(value), int64(0), "created_at 必须是有效的 unix 时间戳")
|
||||
return int64(value)
|
||||
}
|
||||
|
||||
// omitempty 陷阱守卫:created_at 为 0 时也必须出现在线格式里,
|
||||
// 否则「字段存在」这件事就依赖于运行时恰好非零。
|
||||
func TestWire_CreatedAtPresentEvenAtZero(t *testing.T) {
|
||||
resp := responseObjectOf(t, ResponsesStreamEvent{
|
||||
Type: "response.created",
|
||||
Response: &ResponsesResponse{ID: "resp_1", Object: "response", Status: "in_progress"},
|
||||
})
|
||||
require.Contains(t, resp, "created_at", "created_at 不得带 omitempty")
|
||||
require.EqualValues(t, 0, resp["created_at"])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Chat Completions → Responses
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestChatCompletionsResponseToResponses_CarriesCreatedAt(t *testing.T) {
|
||||
t.Run("uses_upstream_created_when_present", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(&ChatCompletionsResponse{
|
||||
ID: "chatcmpl_1",
|
||||
Created: 1700000000,
|
||||
Model: "deepseek-v4-flash",
|
||||
Choices: []ChatChoice{{Message: ChatMessage{Role: "assistant", Content: json.RawMessage(`"hi"`)}}},
|
||||
}, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.EqualValues(t, 1700000000, out.CreatedAt, "上游给了 created 就照搬,不要另起时间")
|
||||
})
|
||||
|
||||
t.Run("stamps_now_when_upstream_omits_created", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(&ChatCompletionsResponse{
|
||||
ID: "chatcmpl_2",
|
||||
Model: "deepseek-v4-flash",
|
||||
Choices: []ChatChoice{{Message: ChatMessage{Role: "assistant", Content: json.RawMessage(`"hi"`)}}},
|
||||
}, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.Greater(t, out.CreatedAt, int64(0))
|
||||
})
|
||||
|
||||
t.Run("nil_upstream_response_still_stamps", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(nil, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.Greater(t, out.CreatedAt, int64(0), "空上游响应也必须产出可解析的对象")
|
||||
})
|
||||
}
|
||||
|
||||
// 同一条流里 response.created 与终止事件必须报同一个 created_at
|
||||
// (官方语义:created_at 是这次 response 的创建时刻,不随事件变化)。
|
||||
func TestChatCompletionsToResponsesStream_CreatedAtStableAcrossEvents(t *testing.T) {
|
||||
state := NewChatCompletionsToResponsesStreamState("deepseek-v4-flash")
|
||||
require.Greater(t, state.Created, int64(0), "前提:state 早就采集了时间戳")
|
||||
|
||||
var chunk ChatCompletionsChunk
|
||||
require.NoError(t, json.Unmarshal(
|
||||
[]byte(`{"choices":[{"index":0,"delta":{"content":"hi"}}]}`), &chunk))
|
||||
|
||||
events := ChatCompletionsChunkToResponsesEvents(&chunk, state)
|
||||
events = append(events, FinalizeChatCompletionsResponsesStream(state)...)
|
||||
|
||||
seen := map[string]int64{}
|
||||
for _, evt := range events {
|
||||
if evt.Response == nil {
|
||||
continue
|
||||
}
|
||||
seen[evt.Type] = requireCreatedAt(t, responseObjectOf(t, evt))
|
||||
}
|
||||
|
||||
require.Contains(t, seen, "response.created")
|
||||
require.Contains(t, seen, "response.completed")
|
||||
require.Equal(t, state.Created, seen["response.created"])
|
||||
require.Equal(t, seen["response.created"], seen["response.completed"],
|
||||
"同一条流的 created_at 必须恒定")
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Anthropic → Responses
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAnthropicToResponsesResponse_StampsCreatedAt(t *testing.T) {
|
||||
out := AnthropicToResponsesResponse(&AnthropicResponse{
|
||||
ID: "msg_1",
|
||||
Type: "message",
|
||||
Role: "assistant",
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
Content: []AnthropicContentBlock{{Type: "text", Text: "hi"}},
|
||||
})
|
||||
require.Greater(t, out.CreatedAt, int64(0),
|
||||
"Anthropic 响应不带时间戳,网关必须自己盖一个")
|
||||
}
|
||||
|
||||
func TestAnthropicEventToResponsesStream_CreatedAtStableAcrossEvents(t *testing.T) {
|
||||
state := NewAnthropicEventToResponsesState()
|
||||
state.Model = "claude-sonnet-4-20250514"
|
||||
require.Greater(t, state.Created, int64(0), "前提:state 早就采集了时间戳")
|
||||
|
||||
var events []ResponsesStreamEvent
|
||||
for _, raw := range []string{
|
||||
`{"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-20250514","content":[],"usage":{"input_tokens":3,"output_tokens":0}}}`,
|
||||
`{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`,
|
||||
`{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}`,
|
||||
`{"type":"content_block_stop","index":0}`,
|
||||
`{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":2}}`,
|
||||
`{"type":"message_stop"}`,
|
||||
} {
|
||||
var evt AnthropicStreamEvent
|
||||
require.NoError(t, json.Unmarshal([]byte(raw), &evt))
|
||||
events = append(events, AnthropicEventToResponsesEvents(&evt, state)...)
|
||||
}
|
||||
events = append(events, FinalizeAnthropicResponsesStream(state)...)
|
||||
|
||||
seen := map[string]int64{}
|
||||
for _, evt := range events {
|
||||
if evt.Response == nil {
|
||||
continue
|
||||
}
|
||||
seen[evt.Type] = requireCreatedAt(t, responseObjectOf(t, evt))
|
||||
}
|
||||
|
||||
require.Contains(t, seen, "response.created")
|
||||
require.Contains(t, seen, "response.completed")
|
||||
require.Equal(t, state.Created, seen["response.created"])
|
||||
require.Equal(t, seen["response.created"], seen["response.completed"],
|
||||
"同一条流的 created_at 必须恒定")
|
||||
}
|
||||
|
||||
// ResponsesClientToolStreamRestorer 对部分事件走 unmarshal→re-marshal。
|
||||
// 结构体没有该字段时,上游带来的 created_at 会在这一步被静默抹掉。
|
||||
func TestResponsesStreamEvent_CreatedAtSurvivesUnmarshalRemarshal(t *testing.T) {
|
||||
upstream := []byte(`{"type":"response.completed","response":{"id":"resp_9","object":"response",` +
|
||||
`"created_at":1700000123,"model":"gpt-5.5","status":"completed","output":[]}}`)
|
||||
|
||||
var evt ResponsesStreamEvent
|
||||
require.NoError(t, json.Unmarshal(upstream, &evt))
|
||||
require.EqualValues(t, 1700000123, evt.Response.CreatedAt)
|
||||
|
||||
require.EqualValues(t, 1700000123, requireCreatedAt(t, responseObjectOf(t, evt)))
|
||||
}
|
||||
@@ -358,8 +358,13 @@ func (t *ResponsesTool) UnmarshalJSON(data []byte) error {
|
||||
|
||||
// ResponsesResponse is the non-streaming response from POST /v1/responses.
|
||||
type ResponsesResponse struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"` // "response"
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"` // "response"
|
||||
// CreatedAt is the unix creation timestamp. Strict Responses clients declare
|
||||
// it non-optional and abort with `missing field 'created_at'` when it is
|
||||
// absent, so it is always emitted — no omitempty. Same rule as ID (see the
|
||||
// "clients treat it as required" fallback in ChatCompletionsResponseToAnthropic).
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Model string `json:"model"`
|
||||
Status string `json:"status"` // "completed" | "incomplete" | "failed"
|
||||
Output []ResponsesOutput `json:"output"`
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -247,18 +247,28 @@ WITH dedup AS (
|
||||
FROM ops_error_logs
|
||||
WHERE created_at >= $1 AND created_at < $2 AND NULLIF(request_id, '') IS NOT NULL
|
||||
)
|
||||
SELECT DISTINCT ON (COALESCE(NULLIF(request_id, ''), 'error:' || id::text))
|
||||
date_trunc('minute', created_at) AS bucket_start,
|
||||
lower(COALESCE(NULLIF(TRIM(platform), ''), 'unknown')) AS platform,
|
||||
COALESCE(group_id, 0) AS group_id,
|
||||
COALESCE(NULLIF(TRIM(requested_model), ''), NULLIF(TRIM(model), ''), 'unknown') AS model,
|
||||
user_id, error_type, error_owner, COALESCE(status_code, 0) AS status_code,
|
||||
COALESCE(upstream_status_code, 0) AS upstream_status_code,
|
||||
lower(CONCAT_WS(' ', error_type, error_source, error_message, upstream_error_message, upstream_error_detail, error_body)) AS text,
|
||||
(CASE WHEN jsonb_typeof(upstream_errors) = 'array' THEN jsonb_array_length(upstream_errors) > 0 ELSE FALSE END
|
||||
OR error_owner = 'provider' OR upstream_status_code IS NOT NULL) AS upstream_affected,
|
||||
CASE WHEN jsonb_typeof(upstream_errors) = 'array' THEN jsonb_array_length(upstream_errors) ELSE 0 END AS upstream_attempts
|
||||
SELECT DISTINCT ON (COALESCE(NULLIF(current_error.request_id, ''), 'error:' || current_error.id::text))
|
||||
date_trunc('minute', current_error.created_at) AS bucket_start,
|
||||
-- Composite groups are a routing layer: resolve the concrete account
|
||||
-- platform (mirrors usageLogEffectivePlatformExpr on the usage side) so
|
||||
-- error facts share the usage facts' platform key. Without this, composite
|
||||
-- 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')
|
||||
ELSE COALESCE(NULLIF(TRIM(current_error.platform), ''), 'unknown')
|
||||
END) AS platform,
|
||||
COALESCE(current_error.group_id, 0) AS group_id,
|
||||
COALESCE(NULLIF(TRIM(current_error.requested_model), ''), NULLIF(TRIM(current_error.model), ''), 'unknown') AS model,
|
||||
current_error.user_id, current_error.error_type, current_error.error_owner, COALESCE(current_error.status_code, 0) AS status_code,
|
||||
COALESCE(current_error.upstream_status_code, 0) AS upstream_status_code,
|
||||
lower(CONCAT_WS(' ', current_error.error_type, current_error.error_source, current_error.error_message, current_error.upstream_error_message, current_error.upstream_error_detail, current_error.error_body)) AS text,
|
||||
(CASE WHEN jsonb_typeof(current_error.upstream_errors) = 'array' THEN jsonb_array_length(current_error.upstream_errors) > 0 ELSE FALSE END
|
||||
OR current_error.error_owner = 'provider' OR current_error.upstream_status_code IS NOT NULL) AS upstream_affected,
|
||||
CASE WHEN jsonb_typeof(current_error.upstream_errors) = 'array' THEN jsonb_array_length(current_error.upstream_errors) ELSE 0 END AS upstream_attempts
|
||||
FROM ops_error_logs current_error
|
||||
LEFT JOIN groups g ON g.id = current_error.group_id
|
||||
LEFT JOIN accounts a ON a.id = current_error.account_id
|
||||
WHERE (
|
||||
(NULLIF(current_error.request_id, '') IS NULL AND current_error.created_at >= $1 AND current_error.created_at < $2)
|
||||
OR (
|
||||
@@ -269,7 +279,7 @@ WITH dedup AS (
|
||||
)
|
||||
AND NOT current_error.is_count_tokens
|
||||
AND (COALESCE(current_error.status_code, 0) >= 400 OR current_error.error_type = 'cyber_policy')
|
||||
ORDER BY COALESCE(NULLIF(request_id, ''), 'error:' || id::text), created_at DESC, id DESC
|
||||
ORDER BY COALESCE(NULLIF(current_error.request_id, ''), 'error:' || current_error.id::text), current_error.created_at DESC, current_error.id DESC
|
||||
), classified AS (
|
||||
SELECT *, CASE
|
||||
-- Keep in lockstep with service.ClassifyChannelMonitorV2Error needles.
|
||||
|
||||
@@ -101,12 +101,25 @@ func TestChannelMonitorV2ErrorAggregationCountsFinalUserErrorsOnly(t *testing.T)
|
||||
require.Contains(t, query, "candidate_ids")
|
||||
require.Contains(t, query, "where bucket_start >= $1 and bucket_start < $2")
|
||||
require.Contains(t, query, "upstream_affected_requests")
|
||||
require.Contains(t, query, "jsonb_array_length(upstream_errors) > 0")
|
||||
require.Contains(t, query, "jsonb_array_length(current_error.upstream_errors) > 0")
|
||||
// request_id dedup must be time-bounded (no full-history scan).
|
||||
require.Contains(t, query, "interval '90 minutes'")
|
||||
require.Contains(t, query, "current_error.created_at >= $1 - interval '90 minutes'")
|
||||
}
|
||||
|
||||
func TestChannelMonitorV2ErrorAggregationResolvesCompositePlatform(t *testing.T) {
|
||||
query := strings.ToLower(channelMonitorV2ErrorAggregationSQL)
|
||||
// Composite groups are a routing layer: error facts must resolve the concrete
|
||||
// account platform (joining groups/accounts) so they aggregate under the same
|
||||
// platform key as usage facts instead of the never-enabled 'composite' platform.
|
||||
require.Contains(t, query, "g.platform = 'composite'")
|
||||
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) {
|
||||
for _, query := range []string{channelMonitorV2UsageMetricsSQL, channelMonitorV2UserMetricsSQL} {
|
||||
require.Contains(t, query, "COALESCE(ul.request_type, 0) NOT IN (4, 6)")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -43,8 +43,9 @@ import (
|
||||
var sseDataPrefix = regexp.MustCompile(`^data:\s*`)
|
||||
|
||||
const (
|
||||
testClaudeAPIURL = "https://api.anthropic.com/v1/messages?beta=true"
|
||||
chatgptCodexAPIURL = "https://chatgpt.com/backend-api/codex/responses"
|
||||
testClaudeAPIURL = "https://api.anthropic.com/v1/messages?beta=true"
|
||||
chatgptCodexAPIURL = "https://chatgpt.com/backend-api/codex/responses"
|
||||
defaultAntigravityTestModel = "claude-sonnet-4-6"
|
||||
)
|
||||
|
||||
// TestEvent represents a SSE event for account testing
|
||||
@@ -146,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
|
||||
@@ -2311,11 +2315,7 @@ func (s *AccountTestService) routeAntigravityTest(c *gin.Context, account *Accou
|
||||
func (s *AccountTestService) testAntigravityAccountConnection(c *gin.Context, account *Account, modelID string) error {
|
||||
ctx := c.Request.Context()
|
||||
|
||||
// 默认模型:Claude 使用 claude-sonnet-4-5,Gemini 使用 gemini-3-pro-preview
|
||||
testModelID := modelID
|
||||
if testModelID == "" {
|
||||
testModelID = "claude-sonnet-4-5"
|
||||
}
|
||||
testModelID := antigravityConnectionTestModel(modelID)
|
||||
|
||||
if s.antigravityGatewayService == nil {
|
||||
return s.sendErrorAndEnd(c, "Antigravity gateway service not configured")
|
||||
@@ -2346,6 +2346,13 @@ func (s *AccountTestService) testAntigravityAccountConnection(c *gin.Context, ac
|
||||
return nil
|
||||
}
|
||||
|
||||
func antigravityConnectionTestModel(modelID string) string {
|
||||
if modelID == "" {
|
||||
return defaultAntigravityTestModel
|
||||
}
|
||||
return modelID
|
||||
}
|
||||
|
||||
// buildGeminiAPIKeyRequest builds request for Gemini API Key accounts
|
||||
func (s *AccountTestService) buildGeminiAPIKeyRequest(ctx context.Context, account *Account, modelID string, payload []byte) (*http.Request, error) {
|
||||
apiKey := account.GetCredential("api_key")
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAntigravityConnectionTestModel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, "claude-sonnet-4-6", antigravityConnectionTestModel(""))
|
||||
require.Equal(t, "gemini-3.1-pro-preview", antigravityConnectionTestModel("gemini-3.1-pro-preview"))
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -69,10 +69,10 @@ func TestAntigravityGatewayService_GetMappedModel(t *testing.T) {
|
||||
expected: "claude-sonnet-4-6",
|
||||
},
|
||||
{
|
||||
name: "默认映射 - claude-sonnet-4-5-20250929 → claude-sonnet-4-5",
|
||||
name: "默认映射 - claude-sonnet-4-5-20250929 → claude-sonnet-4-6",
|
||||
requestedModel: "claude-sonnet-4-5-20250929",
|
||||
accountMapping: nil,
|
||||
expected: "claude-sonnet-4-5",
|
||||
expected: "claude-sonnet-4-6",
|
||||
},
|
||||
|
||||
// 3. 默认映射中的透传(映射到自己)
|
||||
@@ -89,7 +89,7 @@ func TestAntigravityGatewayService_GetMappedModel(t *testing.T) {
|
||||
expected: "claude-sonnet-4-6",
|
||||
},
|
||||
{
|
||||
name: "默认映射透传 - claude-sonnet-4-5",
|
||||
name: "显式 canonical 选择 - claude-sonnet-4-5 透传",
|
||||
requestedModel: "claude-sonnet-4-5",
|
||||
accountMapping: nil,
|
||||
expected: "claude-sonnet-4-5",
|
||||
@@ -113,10 +113,19 @@ func TestAntigravityGatewayService_GetMappedModel(t *testing.T) {
|
||||
expected: "claude-opus-4-6-thinking",
|
||||
},
|
||||
{
|
||||
name: "默认映射透传 - claude-sonnet-4-5-thinking",
|
||||
name: "默认映射 - claude-sonnet-4-5-thinking → claude-sonnet-4-6",
|
||||
requestedModel: "claude-sonnet-4-5-thinking",
|
||||
accountMapping: nil,
|
||||
expected: "claude-sonnet-4-5-thinking",
|
||||
expected: "claude-sonnet-4-6",
|
||||
},
|
||||
{
|
||||
name: "账户显式目标只映射一步 - custom-sonnet → claude-sonnet-4-5",
|
||||
requestedModel: "custom-sonnet",
|
||||
accountMapping: map[string]string{
|
||||
"custom-sonnet": "claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-6",
|
||||
},
|
||||
expected: "claude-sonnet-4-5",
|
||||
},
|
||||
{
|
||||
name: "默认映射透传 - gemini-2.5-flash",
|
||||
|
||||
@@ -844,17 +844,28 @@ func TestSetAntigravityModelRateLimits_GeminiWritesFamilyScope(t *testing.T) {
|
||||
require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
|
||||
}
|
||||
|
||||
func TestSetAntigravityModelRateLimits_ClaudeDoesNotWriteGeminiScope(t *testing.T) {
|
||||
func TestSetAntigravityModelRateLimits_DoesNotDoubleMapCustomChain(t *testing.T) {
|
||||
repo := &stubAntigravityAccountRepo{}
|
||||
svc := &AntigravityGatewayService{}
|
||||
account := &Account{ID: 790, Platform: PlatformAntigravity}
|
||||
account := &Account{
|
||||
ID: 790,
|
||||
Platform: PlatformAntigravity,
|
||||
Credentials: map[string]any{
|
||||
"model_mapping": map[string]any{
|
||||
"custom-sonnet": "claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-6",
|
||||
},
|
||||
},
|
||||
}
|
||||
resetAt := time.Now().Add(30 * time.Second)
|
||||
canonicalModel := resolveFinalAntigravityModelKey(context.Background(), account, "custom-sonnet")
|
||||
require.Equal(t, "claude-sonnet-4-5", canonicalModel)
|
||||
|
||||
success := svc.setAntigravityModelRateLimits(
|
||||
context.Background(),
|
||||
repo,
|
||||
account,
|
||||
"claude-sonnet-4-5",
|
||||
canonicalModel,
|
||||
"[test]",
|
||||
429,
|
||||
resetAt,
|
||||
@@ -866,6 +877,33 @@ func TestSetAntigravityModelRateLimits_ClaudeDoesNotWriteGeminiScope(t *testing.
|
||||
require.Equal(t, "claude-sonnet-4-5", repo.modelRateLimitCalls[0].modelKey)
|
||||
}
|
||||
|
||||
func TestSetModelRateLimitAndClearSession_UsesUpstreamReportedModelMetadata(t *testing.T) {
|
||||
repo := &stubAntigravityAccountRepo{}
|
||||
svc := &AntigravityGatewayService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 791,
|
||||
Platform: PlatformAntigravity,
|
||||
Credentials: map[string]any{
|
||||
"model_mapping": map[string]any{
|
||||
"claude-sonnet-4-5": "claude-sonnet-4-6",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
svc.setModelRateLimitAndClearSession(&handleModelRateLimitParams{
|
||||
ctx: context.Background(),
|
||||
prefix: "[test]",
|
||||
account: account,
|
||||
statusCode: 429,
|
||||
}, &antigravitySmartRetryInfo{
|
||||
RetryDelay: 30 * time.Second,
|
||||
ModelName: "claude-sonnet-4-5",
|
||||
})
|
||||
|
||||
require.Len(t, repo.modelRateLimitCalls, 1)
|
||||
require.Equal(t, "claude-sonnet-4-5", repo.modelRateLimitCalls[0].modelKey)
|
||||
}
|
||||
|
||||
func TestAntigravityRetryLoop_PreCheck_SwitchesWhenRateLimited(t *testing.T) {
|
||||
upstream := &recordingOKUpstream{}
|
||||
account := &Account{
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"math"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -1370,16 +1371,43 @@ func (s *BillingService) computeTokenBreakdown(
|
||||
// multiplier 用于长上下文等场景下的整体价格缩放(普通调用传 1.0 即可)。
|
||||
func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens, price, multiplier float64) float64 {
|
||||
if pricing.SupportsCacheBreakdown && (pricing.CacheCreation5mPrice > 0 || pricing.CacheCreation1hPrice > 0) {
|
||||
if tokens.CacheCreation5mTokens == 0 && tokens.CacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 {
|
||||
cacheCreation5mTokens, cacheCreation1hTokens := normalizeCacheCreationBreakdown(tokens)
|
||||
if cacheCreation5mTokens == 0 && cacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 {
|
||||
// API 未返回 ephemeral 明细,回退到全部按 5m 单价计费
|
||||
return float64(tokens.CacheCreationTokens) * pricing.CacheCreation5mPrice * multiplier
|
||||
}
|
||||
return float64(tokens.CacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier +
|
||||
float64(tokens.CacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier
|
||||
return float64(cacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier +
|
||||
float64(cacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier
|
||||
}
|
||||
return float64(tokens.CacheCreationTokens) * price * multiplier
|
||||
}
|
||||
|
||||
// normalizeCacheCreationBreakdown caps contradictory 5m/1h details at an explicitly
|
||||
// positive aggregate while retaining their reported ratio as closely as integer tokens allow.
|
||||
func normalizeCacheCreationBreakdown(tokens UsageTokens) (int, int) {
|
||||
cacheCreation5mTokens := tokens.CacheCreation5mTokens
|
||||
cacheCreation1hTokens := tokens.CacheCreation1hTokens
|
||||
aggregate := tokens.CacheCreationTokens
|
||||
if cacheCreation5mTokens < 0 {
|
||||
cacheCreation5mTokens = 0
|
||||
}
|
||||
if cacheCreation1hTokens < 0 {
|
||||
cacheCreation1hTokens = 0
|
||||
}
|
||||
if aggregate <= 0 || (cacheCreation5mTokens <= aggregate && cacheCreation1hTokens <= aggregate-cacheCreation5mTokens) {
|
||||
return cacheCreation5mTokens, cacheCreation1hTokens
|
||||
}
|
||||
|
||||
detailTotal := float64(cacheCreation5mTokens) + float64(cacheCreation1hTokens)
|
||||
normalized5mTokens := math.Round(float64(aggregate) * float64(cacheCreation5mTokens) / detailTotal)
|
||||
if normalized5mTokens >= float64(aggregate) {
|
||||
cacheCreation5mTokens = aggregate
|
||||
} else {
|
||||
cacheCreation5mTokens = int(normalized5mTokens)
|
||||
}
|
||||
return cacheCreation5mTokens, aggregate - cacheCreation5mTokens
|
||||
}
|
||||
|
||||
// calculatePerRequestCost 按次/图片计费
|
||||
func (s *BillingService) calculatePerRequestCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) {
|
||||
units := input.UsageUnits
|
||||
|
||||
@@ -1384,6 +1384,122 @@ func TestCalculateCost_SupportsCacheBreakdown(t *testing.T) {
|
||||
require.InDelta(t, expected5m+expected1h, cost.CacheCreationCost, 1e-10)
|
||||
}
|
||||
|
||||
func TestComputeCacheCreationCost_CapsContradictoryBreakdownAtAggregate(t *testing.T) {
|
||||
svc := &BillingService{}
|
||||
pricing := &ModelPricing{
|
||||
SupportsCacheBreakdown: true,
|
||||
CacheCreation5mPrice: 1,
|
||||
CacheCreation1hPrice: 1,
|
||||
}
|
||||
|
||||
tokens := UsageTokens{
|
||||
CacheCreationTokens: 463184,
|
||||
CacheCreation5mTokens: 463184,
|
||||
CacheCreation1hTokens: 463184,
|
||||
}
|
||||
|
||||
cost := svc.computeCacheCreationCost(pricing, tokens, 0, 1)
|
||||
require.Equal(t, float64(tokens.CacheCreationTokens), cost,
|
||||
"billed cache-creation token equivalent must not exceed the positive aggregate")
|
||||
}
|
||||
|
||||
func TestNormalizeCacheCreationBreakdown_BillingSafetyInvariant(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tokens UsageTokens
|
||||
want5m int
|
||||
want1h int
|
||||
}{
|
||||
{
|
||||
name: "preserves ratio when capping",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 90, CacheCreation1hTokens: 60},
|
||||
want5m: 60,
|
||||
want1h: 40,
|
||||
},
|
||||
{
|
||||
name: "details below aggregate unchanged",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 30, CacheCreation1hTokens: 60},
|
||||
want5m: 30,
|
||||
want1h: 60,
|
||||
},
|
||||
{
|
||||
name: "absent 5m detail unchanged",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation1hTokens: 60},
|
||||
want5m: 0,
|
||||
want1h: 60,
|
||||
},
|
||||
{
|
||||
name: "absent 1h detail unchanged",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 30},
|
||||
want5m: 30,
|
||||
want1h: 0,
|
||||
},
|
||||
{
|
||||
name: "negative detail clamped",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -50, CacheCreation1hTokens: 60},
|
||||
want5m: 0,
|
||||
want1h: 60,
|
||||
},
|
||||
{
|
||||
name: "negative detail cannot hide oversized positive detail",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -50, CacheCreation1hTokens: 150},
|
||||
want5m: 0,
|
||||
want1h: 100,
|
||||
},
|
||||
{
|
||||
name: "integer boundary details capped without overflow",
|
||||
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: int(^uint(0) >> 1), CacheCreation1hTokens: int(^uint(0) >> 1)},
|
||||
want5m: 50,
|
||||
want1h: 50,
|
||||
},
|
||||
{
|
||||
name: "integer boundary aggregate avoids float conversion overflow",
|
||||
tokens: UsageTokens{CacheCreationTokens: int(^uint(0) >> 1), CacheCreation5mTokens: int(^uint(0) >> 1), CacheCreation1hTokens: 1},
|
||||
want5m: int(^uint(0) >> 1),
|
||||
want1h: 0,
|
||||
},
|
||||
{
|
||||
name: "zero aggregate unchanged",
|
||||
tokens: UsageTokens{CacheCreation5mTokens: 90, CacheCreation1hTokens: 60},
|
||||
want5m: 90,
|
||||
want1h: 60,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got5m, got1h := normalizeCacheCreationBreakdown(tt.tokens)
|
||||
require.Equal(t, tt.want5m, got5m)
|
||||
require.Equal(t, tt.want1h, got1h)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeCacheCreationCost_PreservesZeroDetailFallback(t *testing.T) {
|
||||
svc := &BillingService{}
|
||||
pricing := &ModelPricing{
|
||||
SupportsCacheBreakdown: true,
|
||||
CacheCreation5mPrice: 4e-6,
|
||||
CacheCreation1hPrice: 5e-6,
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
tokens UsageTokens
|
||||
}{
|
||||
{name: "zero details", tokens: UsageTokens{CacheCreationTokens: 100}},
|
||||
{name: "one negative detail", tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -25}},
|
||||
{name: "both negative details", tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -25, CacheCreation1hTokens: -75}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cost := svc.computeCacheCreationCost(pricing, tt.tokens, 0, 1)
|
||||
require.InDelta(t, 100*4e-6, cost, 1e-12)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateCost_LargeTokenCount(t *testing.T) {
|
||||
svc := newTestBillingService()
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -139,7 +139,9 @@ func DetectModelPlatform(model string) (string, bool) {
|
||||
return PlatformGemini, true
|
||||
case normalized == "grok" || strings.HasPrefix(normalized, "grok-"):
|
||||
return PlatformGrok, true
|
||||
case strings.HasPrefix(normalized, "kimi-"),
|
||||
case normalized == "k3",
|
||||
normalized == "k3-256k",
|
||||
strings.HasPrefix(normalized, "kimi-"),
|
||||
strings.HasPrefix(normalized, "moonshot-"):
|
||||
return PlatformKimi, true
|
||||
case strings.HasPrefix(normalized, "glm-"):
|
||||
|
||||
@@ -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
|
||||
@@ -26,9 +170,13 @@ func TestDetectModelPlatform(t *testing.T) {
|
||||
{name: "grok", model: "grok-4", platform: PlatformGrok, ok: true},
|
||||
{name: "xai prefix", model: "xai/grok-4", platform: PlatformGrok, ok: true},
|
||||
{name: "kimi", model: "kimi-k2-thinking", platform: PlatformKimi, ok: true},
|
||||
{name: "kimi code bare k3", model: "K3", platform: PlatformKimi, ok: true},
|
||||
{name: "kimi code bare k3 256k", model: "k3-256k", platform: PlatformKimi, ok: true},
|
||||
{name: "kimi code provider prefix", model: "kimi-code/k3", platform: PlatformKimi, ok: true},
|
||||
{name: "moonshot prefix", model: "moonshot/moonshot-v1-32k", platform: PlatformKimi, ok: true},
|
||||
{name: "zhipu", model: "glm-5.2", platform: PlatformZhipu, ok: true},
|
||||
{name: "deepseek", model: "deepseek-v4-pro", platform: PlatformDeepseek, ok: true},
|
||||
{name: "unknown k3 alias", model: "k3-preview", ok: false},
|
||||
{name: "unknown", model: "llama-4-maverick", ok: false},
|
||||
}
|
||||
|
||||
|
||||
@@ -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{
|
||||
@@ -191,6 +284,22 @@ func TestCompositeRouteResolverIgnoresDisabledRoutesAndFallsBackToDetector(t *te
|
||||
require.Nil(t, decision.Route)
|
||||
}
|
||||
|
||||
func TestCompositeRouteResolverDetectsKimiCodeBareModels(t *testing.T) {
|
||||
resolver := NewCompositeRouteResolver(nil)
|
||||
|
||||
for _, model := range []string{"k3", "k3-256k", "kimi-code/k3"} {
|
||||
t.Run(model, func(t *testing.T) {
|
||||
decision, err := resolver.Resolve(context.Background(), 7, model, CompositeRouteEndpointMessages)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, decision.Matched)
|
||||
require.Equal(t, CompositeRouteSourceDetector, decision.Source)
|
||||
require.Equal(t, PlatformKimi, decision.TargetPlatform)
|
||||
require.Equal(t, model, decision.UpstreamModel)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeRouteResolverExplicitRoutesCoverBucketTwoProviders(t *testing.T) {
|
||||
resolver := NewCompositeRouteResolver(compositeRouteRepoStub{
|
||||
routes: []CompositeModelRoute{
|
||||
|
||||
@@ -1295,21 +1295,20 @@ func TestGatewayService_ParseSSEUsagePassthrough_MessageStartFallbacks(t *testin
|
||||
}
|
||||
|
||||
func TestGatewayService_ParseSSEUsagePassthrough_MessageDeltaSelectiveOverwrite(t *testing.T) {
|
||||
usage := &ClaudeUsage{
|
||||
InputTokens: 10,
|
||||
CacheCreation5mTokens: 2,
|
||||
CacheCreation1hTokens: 6,
|
||||
}
|
||||
data := `{"type":"message_delta","usage":{"input_tokens":0,"output_tokens":5,"cache_creation_input_tokens":8,"cache_read_input_tokens":0,"cached_tokens":11,"cache_creation":{"ephemeral_5m_input_tokens":1,"ephemeral_1h_input_tokens":0}}}`
|
||||
usage := &ClaudeUsage{}
|
||||
start := `{"type":"message_start","message":{"usage":{"input_tokens":10,"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":463184}}}}`
|
||||
parseSSEUsagePassthrough(start, usage)
|
||||
|
||||
data := `{"type":"message_delta","usage":{"input_tokens":0,"output_tokens":5,"cache_creation_input_tokens":463184,"cache_read_input_tokens":0,"cached_tokens":11,"cache_creation":{"ephemeral_5m_input_tokens":463184,"ephemeral_1h_input_tokens":0}}}`
|
||||
|
||||
parseSSEUsagePassthrough(data, usage)
|
||||
|
||||
require.Equal(t, 10, usage.InputTokens, "message_delta 中 0 值不应覆盖已有 input_tokens")
|
||||
require.Equal(t, 5, usage.OutputTokens)
|
||||
require.Equal(t, 8, usage.CacheCreationInputTokens)
|
||||
require.Equal(t, 463184, usage.CacheCreationInputTokens)
|
||||
require.Equal(t, 11, usage.CacheReadInputTokens, "cache_read_input_tokens 为空时应回退到 cached_tokens")
|
||||
require.Equal(t, 1, usage.CacheCreation5mTokens)
|
||||
require.Equal(t, 6, usage.CacheCreation1hTokens, "message_delta 中 0 值不应覆盖已有 1h 明细")
|
||||
require.Equal(t, 463184, usage.CacheCreation5mTokens)
|
||||
require.Equal(t, 0, usage.CacheCreation1hTokens)
|
||||
}
|
||||
|
||||
func TestGatewayService_ParseSSEUsagePassthrough_NoopCases(t *testing.T) {
|
||||
|
||||
@@ -669,10 +669,10 @@ func parseSSEUsagePassthrough(data string, usage *ClaudeUsage) {
|
||||
|
||||
cc5m := deltaUsage.Get("cache_creation.ephemeral_5m_input_tokens")
|
||||
cc1h := deltaUsage.Get("cache_creation.ephemeral_1h_input_tokens")
|
||||
if cc5m.Exists() && cc5m.Int() > 0 {
|
||||
if cc5m.Exists() {
|
||||
usage.CacheCreation5mTokens = int(cc5m.Int())
|
||||
}
|
||||
if cc1h.Exists() && cc1h.Int() > 0 {
|
||||
if cc1h.Exists() {
|
||||
usage.CacheCreation1hTokens = int(cc1h.Int())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 初始化网关调试日志文件。
|
||||
|
||||
@@ -82,20 +82,19 @@ func TestParseSSEUsage_DeltaOverwritesWithNonZero(t *testing.T) {
|
||||
require.Equal(t, 60, usage.CacheReadInputTokens)
|
||||
}
|
||||
|
||||
func TestParseSSEUsage_DeltaDoesNotResetCacheCreationBreakdown(t *testing.T) {
|
||||
func TestParseSSEUsage_DeltaAuthoritativelyUpdatesCacheCreationBreakdown(t *testing.T) {
|
||||
svc := newMinimalGatewayService()
|
||||
usage := &ClaudeUsage{}
|
||||
|
||||
// 先在 message_start 中写入非零 5m/1h 明细
|
||||
svc.parseSSEUsage(`{"type":"message_start","message":{"usage":{"input_tokens":100,"cache_creation":{"ephemeral_5m_input_tokens":30,"ephemeral_1h_input_tokens":70}}}}`, usage)
|
||||
require.Equal(t, 30, usage.CacheCreation5mTokens)
|
||||
require.Equal(t, 70, usage.CacheCreation1hTokens)
|
||||
svc.parseSSEUsage(`{"type":"message_start","message":{"usage":{"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":463184}}}}`, usage)
|
||||
require.Equal(t, 463184, usage.CacheCreationInputTokens)
|
||||
require.Equal(t, 0, usage.CacheCreation5mTokens)
|
||||
require.Equal(t, 463184, usage.CacheCreation1hTokens)
|
||||
|
||||
// 后续 delta 带默认 0,不应覆盖已有非零值
|
||||
svc.parseSSEUsage(`{"type":"message_delta","usage":{"output_tokens":12,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":0}}}`, usage)
|
||||
require.Equal(t, 30, usage.CacheCreation5mTokens, "delta 的 0 值不应重置 5m 明细")
|
||||
require.Equal(t, 70, usage.CacheCreation1hTokens, "delta 的 0 值不应重置 1h 明细")
|
||||
require.Equal(t, 12, usage.OutputTokens)
|
||||
svc.parseSSEUsage(`{"type":"message_delta","usage":{"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":463184,"ephemeral_1h_input_tokens":0}}}`, usage)
|
||||
require.Equal(t, 463184, usage.CacheCreationInputTokens)
|
||||
require.Equal(t, 463184, usage.CacheCreation5mTokens)
|
||||
require.Equal(t, 0, usage.CacheCreation1hTokens)
|
||||
}
|
||||
|
||||
func TestParseSSEUsage_InvalidJSON(t *testing.T) {
|
||||
|
||||
@@ -1237,11 +1237,11 @@ func (s *GatewayService) extractSSEUsagePatch(event map[string]any) *sseUsagePat
|
||||
patch.hasCacheReadInput = true
|
||||
}
|
||||
if cc, ok := usageObj["cache_creation"].(map[string]any); ok {
|
||||
if v, exists := parseSSEUsageInt(cc["ephemeral_5m_input_tokens"]); exists && v > 0 {
|
||||
if v, exists := parseSSEUsageInt(cc["ephemeral_5m_input_tokens"]); exists {
|
||||
patch.cacheCreation5mTokens = v
|
||||
patch.hasCacheCreation5m = true
|
||||
}
|
||||
if v, exists := parseSSEUsageInt(cc["ephemeral_1h_input_tokens"]); exists && v > 0 {
|
||||
if v, exists := parseSSEUsageInt(cc["ephemeral_1h_input_tokens"]); exists {
|
||||
patch.cacheCreation1hTokens = v
|
||||
patch.hasCacheCreation1h = true
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
@@ -114,6 +146,33 @@ func TestOpenAI429RetryDelayHonorsBoundedRetryAfter(t *testing.T) {
|
||||
require.Equal(t, openAIOAuth429MaxRetryDelay, openAIOAuth429SameAccountRetryDelay(http.Header{"Retry-After": []string{"90"}}, deadline))
|
||||
}
|
||||
|
||||
func TestOpenAI429FastPath_OpenCodeGoUsageLimitUsesMessageResetDuration(t *testing.T) {
|
||||
repo := &rateLimit429AccountRepoStub{}
|
||||
rateLimitService := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
|
||||
svc := &OpenAIGatewayService{rateLimitService: rateLimitService}
|
||||
rateLimitService.SetAccountRuntimeBlocker(svc)
|
||||
account := &Account{ID: 44, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
|
||||
body := []byte(`{"type":"error","error":{"type":"GoUsageLimitError","message":"5-hour usage limit reached. Resets in 4hr 59min. To continue using this model now, enable usage from your available balance: https://opencode.ai/workspace/wrk_test/go"},"metadata":{"workspace":"wrk_test","limitName":"5 hour"}}`)
|
||||
|
||||
before := time.Now()
|
||||
shouldDisable := svc.handleOpenAIAccountUpstreamError(
|
||||
context.Background(),
|
||||
account,
|
||||
http.StatusTooManyRequests,
|
||||
http.Header{},
|
||||
body,
|
||||
)
|
||||
after := time.Now()
|
||||
|
||||
require.False(t, shouldDisable)
|
||||
require.Equal(t, 1, repo.rateLimitCalls)
|
||||
require.Equal(t, account.ID, repo.lastRateLimitID)
|
||||
expectedResetAfter := 4*time.Hour + 59*time.Minute
|
||||
require.False(t, repo.lastRateLimitReset.Before(before.Add(expectedResetAfter-time.Second)))
|
||||
require.False(t, repo.lastRateLimitReset.After(after.Add(expectedResetAfter)))
|
||||
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
}
|
||||
|
||||
// TestOpenAI429FastPath_SkipsSparkShadow 外审第8轮 P1:spark 影子被选中后若 /responses 返回 429,
|
||||
// 不得按 global x-codex-* 信号写内存运行时熔断(否则 spark 被冷却到 global reset、单影子场景无可用账号)。
|
||||
func TestOpenAI429FastPath_SkipsSparkShadow(t *testing.T) {
|
||||
@@ -137,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
@@ -91,7 +91,7 @@ func TestApplyCodexOAuthTransform_ReservedPythonNameIsOAuthOnly(t *testing.T) {
|
||||
require.Equal(t, codexPythonToolAlias, tool["name"])
|
||||
|
||||
apiKeyBody := []byte(`{"type":"response.create","tools":[{"type":"function","name":"python"}]}`)
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(apiKeyBody, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey})
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(apiKeyBody, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, string(apiKeyBody), string(normalized))
|
||||
|
||||
@@ -96,7 +96,7 @@ func TestWebSocketCompatibilityNormalizesTriggerAfterPairedOutputCleanup(t *test
|
||||
body := []byte(`{"type":"response.create","model":"gpt-5.4","input":[{"type":"compaction_trigger"},{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},{"type":"function_call_output","call_id":"call_1","output":"ok"},{"type":"message","role":"user","content":"visible"}]}`)
|
||||
account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, account)
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, account, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
items := gjson.GetBytes(normalized, "input").Array()
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
@@ -117,8 +118,11 @@ func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, errType
|
||||
"response": map[string]any{
|
||||
"id": "resp_" + strings.ReplaceAll(uuid.NewString(), "-", ""),
|
||||
"object": "response",
|
||||
"status": "failed",
|
||||
"output": []any{},
|
||||
// 严格客户端把 created_at 当必填字段,缺失会反序列化失败,
|
||||
// 终止事件就白发了(退化成盲重连)。与 writeResponsesFailedSSE 对齐。
|
||||
"created_at": time.Now().Unix(),
|
||||
"status": "failed",
|
||||
"output": []any{},
|
||||
"error": map[string]any{
|
||||
"code": errType,
|
||||
"message": message,
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5601:严格的 Responses 客户端把 created_at 当必填字段,缺失即
|
||||
// `missing field 'created_at'`。writeOpenAICompactSSEFailureMessage 存在的理由就是
|
||||
// 让 Codex 能把这帧识别成合法终止事件;解析不了就退化回它想避免的盲重连。
|
||||
func TestWriteOpenAICompactSSEFailureMessage_CarriesCreatedAt(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||
|
||||
writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, "upstream_error", "boom")
|
||||
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "event: response.failed")
|
||||
|
||||
_, payload, found := strings.Cut(body, "data: ")
|
||||
require.True(t, found, "SSE 帧必须带 data 行: %q", body)
|
||||
|
||||
var event struct {
|
||||
Type string `json:"type"`
|
||||
Response struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Status string `json:"status"`
|
||||
} `json:"response"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal([]byte(strings.TrimSpace(payload)), &event))
|
||||
|
||||
require.Equal(t, "response.failed", event.Type)
|
||||
require.Equal(t, "response", event.Response.Object)
|
||||
require.Equal(t, "failed", event.Response.Status)
|
||||
require.Greater(t, event.Response.CreatedAt, int64(0),
|
||||
"response.failed 必须带有效的 created_at,否则严格客户端读不出这帧")
|
||||
}
|
||||
@@ -271,7 +271,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyPreservesOpaqueRefere
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: accountType,
|
||||
})
|
||||
}, false)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
@@ -285,7 +285,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyPreservesOpaqueRefere
|
||||
second, changedAgain, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(normalized, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: accountType,
|
||||
})
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changedAgain)
|
||||
require.JSONEq(t, string(normalized), string(second))
|
||||
|
||||
@@ -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)
|
||||
@@ -76,8 +85,9 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
body = reasoningBody
|
||||
}
|
||||
}
|
||||
if account.IsOpenAIOAuthLike() && isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader)) {
|
||||
liteBody, changed, liteErr := normalizeOpenAIResponsesLiteToolsPayload(body)
|
||||
responsesLite := account.IsOpenAI() && isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader))
|
||||
if responsesLite {
|
||||
liteBody, changed, liteErr := normalizeOpenAIResponsesLitePayloadForAccount(body, account)
|
||||
if liteErr != nil {
|
||||
param := "tools"
|
||||
var validationErr *openAIResponsesLiteValidationError
|
||||
@@ -111,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 {
|
||||
@@ -151,7 +161,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
return s.forwardResponsesViaNativeAnthropic(ctx, c, account, body, reqModel)
|
||||
}
|
||||
if account.IsOpenAIApiKey() {
|
||||
if normalized, changed, normalizeErr := normalizeOpenAIParallelToolCallsWithoutTools(body); normalizeErr != nil {
|
||||
if normalized, changed, normalizeErr := normalizeOpenAIParallelToolCallsWithoutTools(body, responsesLite); normalizeErr != nil {
|
||||
return nil, normalizeErr
|
||||
} else if changed {
|
||||
body = normalized
|
||||
@@ -172,6 +182,17 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body)
|
||||
}
|
||||
if account.IsOpenAI() && (account.IsOpenAIApiKey() || account.IsOpenAIOAuthLike()) {
|
||||
normalizedReasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningContentReplay(body)
|
||||
if reasoningErr != nil {
|
||||
return nil, fmt.Errorf("normalize OpenAI Responses reasoning content replay: %w", reasoningErr)
|
||||
}
|
||||
if reasoningChanged {
|
||||
body = normalizedReasoningBody
|
||||
originalBody = normalizedReasoningBody
|
||||
requestView = newOpenAIRequestView(normalizedReasoningBody)
|
||||
reqModel, reqStream, promptCacheKey = requestView.Model, requestView.Stream, requestView.PromptCacheKey
|
||||
originalModel = reqModel
|
||||
}
|
||||
sanitizedBody, changed, sanitizeErr := sanitizeOpenAIResponsesInputItemIDs(body)
|
||||
if sanitizeErr != nil {
|
||||
return nil, fmt.Errorf("sanitize OpenAI Responses input item IDs: %w", sanitizeErr)
|
||||
@@ -462,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()}})
|
||||
|
||||
@@ -723,7 +723,7 @@ func normalizeGrokReasoningEffortValue(raw, model string) (string, bool) {
|
||||
case "minimal":
|
||||
return "low", true
|
||||
case "xhigh", "extrahigh":
|
||||
if grokSupportsXHighReasoningEffort(model) {
|
||||
if GrokSupportsXHighReasoningEffort(model) {
|
||||
return "xhigh", true
|
||||
}
|
||||
return "high", true
|
||||
@@ -734,7 +734,9 @@ func normalizeGrokReasoningEffortValue(raw, model string) (string, bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func grokSupportsXHighReasoningEffort(model string) bool {
|
||||
// GrokSupportsXHighReasoningEffort reports whether the model advertises and
|
||||
// forwards the xhigh reasoning effort (Grok 4.6 and its undated alias).
|
||||
func GrokSupportsXHighReasoningEffort(model string) bool {
|
||||
model = strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model)))
|
||||
return model == "grok-4.6" || model == "grok-4.6-latest"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -212,7 +212,8 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
}
|
||||
}
|
||||
if account != nil && account.IsOpenAI() {
|
||||
normalizedBody, normalized, normalizeErr := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, account)
|
||||
responsesLite := isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader)) || isOpenAIResponsesLiteWebSocketPayload(body)
|
||||
normalizedBody, normalized, normalizeErr := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, account, responsesLite)
|
||||
if normalizeErr != nil {
|
||||
return nil, fmt.Errorf("normalize passthrough Responses compatibility: %w", normalizeErr)
|
||||
}
|
||||
@@ -1677,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 协议账号原样返回。
|
||||
@@ -349,7 +412,7 @@ func normalizeOpenAICompactRequestBody(body []byte) ([]byte, bool, error) {
|
||||
}
|
||||
normalized = next
|
||||
}
|
||||
if next, removed, err := normalizeOpenAIParallelToolCallsWithoutTools(normalized); err != nil {
|
||||
if next, removed, err := normalizeOpenAIParallelToolCallsWithoutTools(normalized, false); err != nil {
|
||||
return body, false, err
|
||||
} else if removed {
|
||||
normalized = next
|
||||
@@ -361,7 +424,10 @@ func normalizeOpenAICompactRequestBody(body []byte) ([]byte, bool, error) {
|
||||
return normalized, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIParallelToolCallsWithoutTools(body []byte) ([]byte, bool, error) {
|
||||
func normalizeOpenAIParallelToolCallsWithoutTools(body []byte, responsesLite bool) ([]byte, bool, error) {
|
||||
if responsesLite {
|
||||
return body, false, nil
|
||||
}
|
||||
parallel := gjson.GetBytes(body, "parallel_tool_calls")
|
||||
if !parallel.Exists() {
|
||||
return body, false, nil
|
||||
@@ -376,16 +442,7 @@ func normalizeOpenAIParallelToolCallsWithoutTools(body []byte) ([]byte, bool, er
|
||||
return normalized, true, nil
|
||||
}
|
||||
|
||||
// openAIRequestBodyHasTools is the []byte counterpart of openAIResponsesLiteHasTools:
|
||||
// besides the top-level "tools" array it also recognizes the Responses Lite carrier.
|
||||
// normalizeOpenAIResponsesLiteTools moves namespace tools into an input item of type
|
||||
// "additional_tools" and drops the top-level "tools" key; the request still carries
|
||||
// tools at that point. Looking only at the top level therefore misreads such a body as
|
||||
// "no tools" and deletes the parallel_tool_calls:false that
|
||||
// ensureOpenAIResponsesLiteParallelToolCalls had just pinned, and OpenAI falls back to
|
||||
// its default of true and rejects the request with
|
||||
// 400 unsupported_value: "X-OpenAI-Internal-Codex-Responses-Lite requires
|
||||
// `parallel_tool_calls` to be false."
|
||||
// openAIRequestBodyHasTools 同时识别顶层 tools 和 input[].additional_tools。
|
||||
func openAIRequestBodyHasTools(body []byte) bool {
|
||||
if tools := gjson.GetBytes(body, "tools"); tools.IsArray() && len(tools.Array()) > 0 {
|
||||
return true
|
||||
@@ -401,6 +458,67 @@ func openAIRequestBodyHasTools(body []byte) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// normalizeOpenAIResponsesReasoningContentReplay removes non-portable
|
||||
// reasoning.content arrays before history is sent to a real OpenAI Responses
|
||||
// endpoint. Compatible providers may return visible reasoning blocks there,
|
||||
// while OpenAI accepts only an empty array when the item is replayed.
|
||||
//
|
||||
// Keep the reasoning item and its portable fields (summary, encrypted_content,
|
||||
// ids, and opaque extensions). Callers scope this normalization to OpenAI
|
||||
// destinations; compatible providers may still consume their own content.
|
||||
func normalizeOpenAIResponsesReasoningContentReplay(body []byte) ([]byte, bool, error) {
|
||||
input := gjson.GetBytes(body, "input")
|
||||
if !input.IsArray() {
|
||||
return body, false, nil
|
||||
}
|
||||
|
||||
needsNormalization := false
|
||||
input.ForEach(func(_, item gjson.Result) bool {
|
||||
if strings.TrimSpace(item.Get("type").String()) != "reasoning" {
|
||||
return true
|
||||
}
|
||||
content := item.Get("content")
|
||||
if content.IsArray() && len(content.Array()) > 0 {
|
||||
needsNormalization = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})
|
||||
if !needsNormalization {
|
||||
return body, false, nil
|
||||
}
|
||||
|
||||
var reqBody map[string]any
|
||||
if err := decodeOpenAIJSONUseNumber(body, &reqBody); err != nil {
|
||||
return body, false, fmt.Errorf("normalize OpenAI reasoning content replay: %w", err)
|
||||
}
|
||||
items, ok := reqBody["input"].([]any)
|
||||
if !ok {
|
||||
return body, false, nil
|
||||
}
|
||||
changed := false
|
||||
for _, rawItem := range items {
|
||||
item, ok := rawItem.(map[string]any)
|
||||
if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) != "reasoning" {
|
||||
continue
|
||||
}
|
||||
content, ok := item["content"].([]any)
|
||||
if !ok || len(content) == 0 {
|
||||
continue
|
||||
}
|
||||
delete(item, "content")
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
return body, false, nil
|
||||
}
|
||||
normalized, err := marshalOpenAIUpstreamJSON(reqBody)
|
||||
if err != nil {
|
||||
return body, false, fmt.Errorf("serialize normalized OpenAI reasoning content replay: %w", err)
|
||||
}
|
||||
return normalized, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIAPIKeyStoreFalseReasoningReplay(body []byte, knownStoreFalse bool) ([]byte, bool, error) {
|
||||
if !knownStoreFalse && gjson.GetBytes(body, "store").Type != gjson.False {
|
||||
return body, false, nil
|
||||
@@ -948,7 +1066,7 @@ func normalizeOpenAIResponseFormatSchemasBody(body []byte) ([]byte, bool, error)
|
||||
return normalized, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Account) ([]byte, bool, error) {
|
||||
func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Account, responsesLite bool) ([]byte, bool, error) {
|
||||
if account == nil || !account.IsOpenAI() {
|
||||
return body, false, nil
|
||||
}
|
||||
@@ -961,8 +1079,14 @@ func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Ac
|
||||
return body, false, err
|
||||
}
|
||||
}
|
||||
if next, normalizedReasoningContent, err := normalizeOpenAIResponsesReasoningContentReplay(normalized); err != nil {
|
||||
return body, false, err
|
||||
} else if normalizedReasoningContent {
|
||||
normalized = next
|
||||
changed = true
|
||||
}
|
||||
if account.IsOpenAIApiKey() {
|
||||
if next, normalizedParallel, err := normalizeOpenAIParallelToolCallsWithoutTools(normalized); err != nil {
|
||||
if next, normalizedParallel, err := normalizeOpenAIParallelToolCallsWithoutTools(normalized, responsesLite); err != nil {
|
||||
return body, false, err
|
||||
} else if normalizedParallel {
|
||||
normalized = next
|
||||
|
||||
@@ -259,34 +259,154 @@ func TestNormalizeOpenAIAPIKeyStoreFalseReasoningReplayRejectsEmptyEncryptedCont
|
||||
|
||||
func TestNormalizeOpenAIParallelToolCallsWithoutTools(t *testing.T) {
|
||||
withTools := []byte(`{"tools":[{"type":"function","name":"lookup"}],"parallel_tool_calls":false}`)
|
||||
normalized, changed, err := normalizeOpenAIParallelToolCallsWithoutTools(withTools)
|
||||
normalized, changed, err := normalizeOpenAIParallelToolCallsWithoutTools(withTools, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, string(withTools), string(normalized))
|
||||
|
||||
withoutTools := []byte(`{"input":"hi","parallel_tool_calls":true}`)
|
||||
normalized, changed, err = normalizeOpenAIParallelToolCallsWithoutTools(withoutTools)
|
||||
normalized, changed, err = normalizeOpenAIParallelToolCallsWithoutTools(withoutTools, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.False(t, gjson.GetBytes(normalized, "parallel_tool_calls").Exists())
|
||||
}
|
||||
|
||||
// A Responses Lite body that has already been through normalizeOpenAIResponsesLiteTools
|
||||
// carries its tools in an input item of type "additional_tools" and no longer has a
|
||||
// top-level "tools" key. It still has tools, so the parallel_tool_calls:false that
|
||||
// ensureOpenAIResponsesLiteParallelToolCalls pinned must survive this normalization —
|
||||
// otherwise OpenAI applies its default of true and rejects the request.
|
||||
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}`)
|
||||
normalized, changed, err := normalizeOpenAIParallelToolCallsWithoutTools(liteBody)
|
||||
normalized, changed, err := normalizeOpenAIParallelToolCallsWithoutTools(liteBody, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, gjson.False, gjson.GetBytes(normalized, "parallel_tool_calls").Type)
|
||||
|
||||
// An empty additional_tools item carries no tools, so the field is still dropped.
|
||||
// 非 Lite 请求的空 additional_tools 不构成有效工具声明,字段仍需删除。
|
||||
emptyLiteBody := []byte(`{"input":[{"type":"additional_tools","tools":[]}],"parallel_tool_calls":true}`)
|
||||
normalized, changed, err = normalizeOpenAIParallelToolCallsWithoutTools(emptyLiteBody)
|
||||
normalized, changed, err = normalizeOpenAIParallelToolCallsWithoutTools(emptyLiteBody, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.False(t, gjson.GetBytes(normalized, "parallel_tool_calls").Exists())
|
||||
|
||||
// Lite 请求即使没有工具,也必须保留已经固定的 false。
|
||||
toolLessLiteBody := []byte(`{"input":"hi","parallel_tool_calls":false}`)
|
||||
normalized, changed, err = normalizeOpenAIParallelToolCallsWithoutTools(toolLessLiteBody, true)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, gjson.False, gjson.GetBytes(normalized, "parallel_tool_calls").Type)
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesReasoningContentReplayStripsCrossProviderArray(t *testing.T) {
|
||||
body := []byte(`{"model":"gpt-5.6-sol","input":[` +
|
||||
`{"type":"message","role":"user","content":"one"},` +
|
||||
`{"type":"message","role":"assistant","content":"two"},` +
|
||||
`{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},` +
|
||||
`{"type":"function_call_output","call_id":"call_1","output":"ok"},` +
|
||||
`{"type":"message","role":"user","content":"five"},` +
|
||||
`{"type":"reasoning","id":"rs_provider","summary":[{"type":"summary_text","text":"portable"}],"content":[{"type":"reasoning_text","text":"visible reasoning"}],"opaque":9007199254740993},` +
|
||||
`{"type":"message","role":"assistant","content":[{"type":"output_text","text":"answer"}]}` +
|
||||
`]}`)
|
||||
|
||||
normalized, changed, err := normalizeOpenAIResponsesReasoningContentReplay(body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(normalized, "input.5.type").String())
|
||||
require.False(t, gjson.GetBytes(normalized, "input.5.content").Exists())
|
||||
require.Equal(t, "portable", gjson.GetBytes(normalized, "input.5.summary.0.text").String())
|
||||
require.Equal(t, "9007199254740993", gjson.GetBytes(normalized, "input.5.opaque").Raw)
|
||||
require.Equal(t, "answer", gjson.GetBytes(normalized, "input.6.content.0.text").String())
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesReasoningContentReplayKeepsPortableShapes(t *testing.T) {
|
||||
for _, body := range []string{
|
||||
`{"input":[{"type":"reasoning","summary":[]}]}`,
|
||||
`{"input":[{"type":"reasoning","content":[],"summary":[]}]}`,
|
||||
`{"input":[{"type":"message","content":[{"type":"input_text","text":"keep"}]}]}`,
|
||||
} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesReasoningContentReplay([]byte(body))
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, body, string(normalized))
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyStripsReasoningContentOnlyForOpenAI(t *testing.T) {
|
||||
body := []byte(`{"type":"response.create","model":"gpt-5.6-sol","store":true,"input":[{"type":"reasoning","summary":[{"type":"summary_text","text":"keep"}],"content":[{"type":"reasoning_text","text":"remove"}]}]}`)
|
||||
for _, accountType := range []string{AccountTypeAPIKey, AccountTypeOAuth} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: accountType,
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.False(t, gjson.GetBytes(normalized, "input.0.content").Exists())
|
||||
require.Equal(t, "keep", gjson.GetBytes(normalized, "input.0.summary.0.text").String())
|
||||
}
|
||||
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformZhipu,
|
||||
Type: AccountTypeAPIKey,
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, string(body), string(normalized))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -250,6 +250,34 @@ func TestOpenAIPassthroughAPIKeyRestoresClientToolsNonStreaming(t *testing.T) {
|
||||
require.Equal(t, "*** Begin Patch", gjson.Get(recorder.Body.String(), "output.1.input").String())
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyPreservesCustomToolOutputContentParts(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"custom","name":"exec"}],"input":[{"type":"custom_tool_call_output","call_id":"call_1","output":[{"type":"input_text","text":"result"},{"type":"input_file","file_id":"file_123"}]}]}`)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_tools","status":"completed","output":[],"usage":{}}`)),
|
||||
}}
|
||||
svc := openAIClientToolsTestService(upstream)
|
||||
account := &Account{ID: 6240, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}}
|
||||
|
||||
result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, false, time.Now())
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.0.type").String())
|
||||
output := gjson.GetBytes(upstream.lastBody, "input.0.output")
|
||||
require.True(t, output.IsArray(), "native Responses content parts must reach the upstream as an array")
|
||||
require.Equal(t, "input_text", output.Get("0.type").String())
|
||||
require.Equal(t, "result", output.Get("0.text").String())
|
||||
require.Equal(t, "input_file", output.Get("1.type").String())
|
||||
require.Equal(t, "file_123", output.Get("1.file_id").String())
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyRestoresClientToolsStreaming(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := openAIClientToolsRequest(true)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -34,12 +34,13 @@ const (
|
||||
openAIImagesGenerationsURL = "https://api.openai.com/v1/images/generations"
|
||||
openAIImagesEditsURL = "https://api.openai.com/v1/images/edits"
|
||||
|
||||
openAIChatGPTStartURL = "https://chatgpt.com/"
|
||||
openAIChatGPTFilesURL = "https://chatgpt.com/backend-api/files"
|
||||
openAIImageBackendUserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
|
||||
openAIImageMaxDownloadBytes = 20 << 20 // 20MB per image download
|
||||
openAIImageMaxUploadPartSize = 20 << 20 // 20MB per multipart upload part
|
||||
openAIImagesResponsesMainModel = "gpt-5.4-mini"
|
||||
openAIChatGPTStartURL = "https://chatgpt.com/"
|
||||
openAIChatGPTFilesURL = "https://chatgpt.com/backend-api/files"
|
||||
openAIImageBackendUserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
|
||||
openAIImageMaxDownloadBytes = 20 << 20 // 20MB per image download
|
||||
openAIImageMaxUploadPartSize = 20 << 20 // 20MB per multipart upload part
|
||||
openAIImagesResponsesMainModel = "gpt-5.4-mini"
|
||||
openAIImagesVerbatimPromptInstructions = "When invoking the image_generation tool, use the user's image prompt verbatim. Do not rewrite, expand, summarize, embellish, translate, normalize punctuation, or add or remove visual details or constraints. Preserve the original language, wording, capitalization, quotes, and punctuation exactly."
|
||||
)
|
||||
|
||||
type OpenAIImagesCapability string
|
||||
|
||||
@@ -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")
|
||||
@@ -356,6 +383,7 @@ func buildOpenAIImagesResponsesRequest(parsed *OpenAIImagesRequest, toolModel st
|
||||
|
||||
req := []byte(`{"instructions":"","stream":true,"reasoning":{"effort":"medium","summary":"auto"},"parallel_tool_calls":true,"include":["reasoning.encrypted_content"],"model":"","store":false,"tool_choice":{"type":"image_generation"}}`)
|
||||
req, _ = sjson.SetBytes(req, "model", openAIImagesResponsesMainModel)
|
||||
req, _ = sjson.SetBytes(req, "instructions", openAIImagesVerbatimPromptInstructions)
|
||||
|
||||
input := []byte(`[{"type":"message","role":"user","content":[{"type":"input_text","text":""}]}]`)
|
||||
input, _ = sjson.SetBytes(input, "0.content.0.text", prompt)
|
||||
@@ -710,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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1774,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
|
||||
@@ -1921,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
|
||||
@@ -2016,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
|
||||
}
|
||||
|
||||
@@ -2011,6 +2011,21 @@ func TestBuildOpenAIImagesResponsesRequest_StripsInputFidelity(t *testing.T) {
|
||||
require.Equal(t, "edit", gjson.GetBytes(body, "tools.0.action").String())
|
||||
}
|
||||
|
||||
func TestBuildOpenAIImagesResponsesRequest_RequiresVerbatimUserPrompt(t *testing.T) {
|
||||
prompt := "画一个蓝色马克杯,杯身只写“SkelOT”,保持大小写;白色背景,不要增加其他文字。"
|
||||
parsed := &OpenAIImagesRequest{
|
||||
Endpoint: openAIImagesGenerationsEndpoint,
|
||||
Model: "gpt-image-2",
|
||||
Prompt: prompt,
|
||||
N: 1,
|
||||
}
|
||||
|
||||
body, err := buildOpenAIImagesResponsesRequest(parsed, "gpt-image-2")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, openAIImagesVerbatimPromptInstructions, gjson.GetBytes(body, "instructions").String())
|
||||
require.Equal(t, prompt, gjson.GetBytes(body, "input.0.content.0.text").String())
|
||||
}
|
||||
|
||||
func TestCollectOpenAIImagesFromResponsesBody_FallsBackToOutputItemDone(t *testing.T) {
|
||||
body := []byte(
|
||||
"data: {\"type\":\"response.created\",\"response\":{\"created_at\":1710000004}}\n\n" +
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -55,7 +55,7 @@ func TestNormalizeOpenAIOAuthResponsesCompatibilityBody_PreservesExplicitInput(t
|
||||
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_OnlyStripsOAuthFields(t *testing.T) {
|
||||
body := []byte(`{"type":"response.create","prompt":"hello","commands":{},"truncation":"auto","stop_sequences":["END"],"chat_template_kwargs":{"enable_thinking":true}}`)
|
||||
|
||||
oauthBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth})
|
||||
oauthBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "hello", gjson.GetBytes(oauthBody, "input").String())
|
||||
@@ -63,7 +63,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_OnlyStripsOAuthField
|
||||
require.False(t, gjson.GetBytes(oauthBody, field).Exists(), field)
|
||||
}
|
||||
|
||||
apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey})
|
||||
apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, string(body), string(apiKeyBody))
|
||||
@@ -81,7 +81,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_SanitizesNativeItemI
|
||||
if oauth {
|
||||
accountType = AccountTypeOAuth
|
||||
}
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType})
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "response.create", gjson.GetBytes(normalized, "type").String())
|
||||
@@ -105,7 +105,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_APIKeyStoreFalseRepl
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
})
|
||||
}, false)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
@@ -141,13 +141,13 @@ func TestNormalizeOpenAIResponsesReasoningMode(t *testing.T) {
|
||||
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_ReasoningModeAccountScope(t *testing.T) {
|
||||
body := []byte(`{"type":"response.create","reasoning":{"mode":"pro"}}`)
|
||||
for _, accountType := range []string{AccountTypeOAuth, AccountTypeSetupToken} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType})
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String())
|
||||
require.False(t, gjson.GetBytes(normalized, "reasoning.mode").Exists())
|
||||
}
|
||||
apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey})
|
||||
apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, string(body), string(apiKeyBody))
|
||||
@@ -156,7 +156,7 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_ReasoningModeAccount
|
||||
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_SanitizesToolSchemas(t *testing.T) {
|
||||
body := []byte(`{"type":"response.create","tools":[{"type":"function","name":"search","parameters":{"type":null,"properties":{"q":{"type":"string","pattern":"^(?=.*foo).+$"}}}}]}`)
|
||||
for _, accountType := range []string{AccountTypeAPIKey, AccountTypeOAuth} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType})
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "object", gjson.GetBytes(normalized, "tools.0.parameters.type").String())
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
@@ -101,9 +100,11 @@ func normalizeOpenAIResponsesLiteTools(reqBody map[string]any) (bool, error) {
|
||||
}
|
||||
|
||||
func ensureOpenAIResponsesLiteParallelToolCalls(reqBody map[string]any, changed bool) (bool, error) {
|
||||
parallel := reqBody["parallel_tool_calls"]
|
||||
if !openAIResponsesLiteHasTools(reqBody) {
|
||||
return changed, nil
|
||||
parallel, exists := reqBody["parallel_tool_calls"]
|
||||
if exists {
|
||||
if _, ok := parallel.(bool); !ok {
|
||||
return false, newOpenAIResponsesLiteValidationError("parallel_tool_calls", "responses Lite requires parallel_tool_calls to be a boolean")
|
||||
}
|
||||
}
|
||||
if parallel == false {
|
||||
return changed, nil
|
||||
@@ -112,23 +113,6 @@ func ensureOpenAIResponsesLiteParallelToolCalls(reqBody map[string]any, changed
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func openAIResponsesLiteHasTools(reqBody map[string]any) bool {
|
||||
if tools, ok := reqBody["tools"].([]any); ok && len(tools) > 0 {
|
||||
return true
|
||||
}
|
||||
input, _ := reqBody["input"].([]any)
|
||||
for _, rawItem := range input {
|
||||
item, ok := rawItem.(map[string]any)
|
||||
if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) != "additional_tools" {
|
||||
continue
|
||||
}
|
||||
if tools, ok := item["tools"].([]any); ok && len(tools) > 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func ensureOpenAIResponsesLiteReasoningContext(reqBody map[string]any) (bool, error) {
|
||||
rawReasoning, exists := reqBody["reasoning"]
|
||||
if !exists || rawReasoning == nil {
|
||||
@@ -254,7 +238,7 @@ func openAIResponsesLiteToolIdentityForError(rawTool any) string {
|
||||
|
||||
func normalizeOpenAIResponsesLiteToolsPayload(body []byte) ([]byte, bool, error) {
|
||||
var requestBody map[string]any
|
||||
if err := json.Unmarshal(body, &requestBody); err != nil {
|
||||
if err := decodeOpenAIJSONUseNumber(body, &requestBody); err != nil {
|
||||
return body, false, fmt.Errorf("decode responses Lite request body: %w", err)
|
||||
}
|
||||
changed, err := normalizeOpenAIResponsesLiteTools(requestBody)
|
||||
@@ -267,3 +251,29 @@ func normalizeOpenAIResponsesLiteToolsPayload(body []byte) ([]byte, bool, error)
|
||||
}
|
||||
return rebuilt, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIResponsesLiteParallelToolCallsPayload(body []byte) ([]byte, bool, error) {
|
||||
var requestBody map[string]any
|
||||
if err := decodeOpenAIJSONUseNumber(body, &requestBody); err != nil {
|
||||
return body, false, fmt.Errorf("decode responses Lite request body: %w", err)
|
||||
}
|
||||
changed, err := ensureOpenAIResponsesLiteParallelToolCalls(requestBody, false)
|
||||
if err != nil || !changed {
|
||||
return body, false, err
|
||||
}
|
||||
rebuilt, err := marshalOpenAIUpstreamJSON(requestBody)
|
||||
if err != nil {
|
||||
return body, false, fmt.Errorf("encode responses Lite request body: %w", err)
|
||||
}
|
||||
return rebuilt, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIResponsesLitePayloadForAccount(body []byte, account *Account) ([]byte, bool, error) {
|
||||
if account == nil || !account.IsOpenAI() {
|
||||
return body, false, nil
|
||||
}
|
||||
if account.IsOpenAIOAuthLike() {
|
||||
return normalizeOpenAIResponsesLiteToolsPayload(body)
|
||||
}
|
||||
return normalizeOpenAIResponsesLiteParallelToolCallsPayload(body)
|
||||
}
|
||||
|
||||
@@ -201,17 +201,33 @@ func TestNormalizeOpenAIResponsesLiteTools_ForcesParallelToolCallsFalse(t *testi
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesLiteTools_DoesNotAddParallelToolCallsWithoutTools(t *testing.T) {
|
||||
reqBody := map[string]any{
|
||||
"reasoning": map[string]any{"context": "all_turns"},
|
||||
"parallel_tool_calls": true,
|
||||
func TestNormalizeOpenAIResponsesLiteTools_PinsParallelToolCallsWithoutTools(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
parallel any
|
||||
include bool
|
||||
wantChanged bool
|
||||
}{
|
||||
{name: "字段缺失", wantChanged: true},
|
||||
{name: "值为 true", parallel: true, include: true, wantChanged: true},
|
||||
{name: "值为 false", parallel: false, include: true, wantChanged: false},
|
||||
}
|
||||
|
||||
changed, err := normalizeOpenAIResponsesLiteTools(reqBody)
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
reqBody := map[string]any{"reasoning": map[string]any{"context": "all_turns"}}
|
||||
if tt.include {
|
||||
reqBody["parallel_tool_calls"] = tt.parallel
|
||||
}
|
||||
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, true, reqBody["parallel_tool_calls"])
|
||||
changed, err := normalizeOpenAIResponsesLiteTools(reqBody)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.wantChanged, changed)
|
||||
require.Contains(t, reqBody, "parallel_tool_calls")
|
||||
require.Equal(t, false, reqBody["parallel_tool_calls"])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesLiteTools_RejectsNonBooleanParallelToolCalls(t *testing.T) {
|
||||
@@ -336,6 +352,34 @@ func TestNormalizeOpenAIResponsesLiteToolsPayload_PreservesResponseCreateShape(t
|
||||
require.False(t, gjson.GetBytes(updated, "parallel_tool_calls").Bool())
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesLitePayloads_PreserveLargeSequence(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"type":"response.create",
|
||||
"sequence":900719925474099312345,
|
||||
"tools":[{"type":"function","name":"lookup"}],
|
||||
"parallel_tool_calls":true
|
||||
}`)
|
||||
tests := []struct {
|
||||
name string
|
||||
normalize func([]byte) ([]byte, bool, error)
|
||||
}{
|
||||
{name: "OAuth-like tools normalization", normalize: normalizeOpenAIResponsesLiteToolsPayload},
|
||||
{name: "API key parallel normalization", normalize: normalizeOpenAIResponsesLiteParallelToolCallsPayload},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
updated, changed, err := tt.normalize(body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "900719925474099312345", gjson.GetBytes(updated, "sequence").Raw)
|
||||
require.True(t, gjson.GetBytes(updated, "parallel_tool_calls").Exists())
|
||||
require.False(t, gjson.GetBytes(updated, "parallel_tool_calls").Bool())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyCodexOAuthTransform_PreservesLiteNamespaceToolChoice(t *testing.T) {
|
||||
reqBody := map[string]any{
|
||||
"model": "gpt-5.6-terra",
|
||||
@@ -453,3 +497,118 @@ func TestOpenAIGatewayServiceForward_NormalizesResponsesLiteToolsForOAuth(t *tes
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceForward_PinsParallelToolCallsForToollessResponsesLite(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
accountCases := []struct {
|
||||
name string
|
||||
accountType string
|
||||
credentials map[string]any
|
||||
}{
|
||||
{name: "oauth", accountType: AccountTypeOAuth, credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-account"}},
|
||||
{name: "apikey", accountType: AccountTypeAPIKey, credentials: map[string]any{"api_key": "sk-test"}},
|
||||
}
|
||||
parallelCases := []struct {
|
||||
name string
|
||||
field string
|
||||
}{
|
||||
{name: "字段缺失"},
|
||||
{name: "值为 true", field: `,"parallel_tool_calls":true`},
|
||||
{name: "值为 false", field: `,"parallel_tool_calls":false`},
|
||||
}
|
||||
|
||||
for _, accountCase := range accountCases {
|
||||
for _, passthrough := range []bool{false, true} {
|
||||
mode := "managed"
|
||||
if passthrough {
|
||||
mode = "passthrough"
|
||||
}
|
||||
for _, parallelCase := range parallelCases {
|
||||
name := accountCase.name + "/" + mode + "/" + parallelCase.name
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
|
||||
c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1")
|
||||
c.Request.Header.Set(responsesLiteHeader, "true")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_lite\",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" +
|
||||
"data: [DONE]\n\n",
|
||||
)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
|
||||
account := &Account{
|
||||
ID: 502, Name: "responses-lite-no-tools", Platform: PlatformOpenAI, Type: accountCase.accountType,
|
||||
Concurrency: 1, Status: StatusActive, Schedulable: true, RateMultiplier: f64p(1),
|
||||
Credentials: accountCase.credentials,
|
||||
Extra: map[string]any{"openai_passthrough": passthrough},
|
||||
}
|
||||
body := []byte(`{
|
||||
"model":"gpt-5.6-terra","stream":true,"instructions":"test",
|
||||
"reasoning":{"effort":"high","context":"current_turn"},
|
||||
"input":[{"type":"message","role":"user","content":"hello"}]` + parallelCase.field + `
|
||||
}`)
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, "true", upstream.lastReq.Header.Get(responsesLiteHeader))
|
||||
require.Equal(t, gjson.False, gjson.GetBytes(upstream.lastBody, "parallel_tool_calls").Type, string(upstream.lastBody))
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceForward_DisablesParallelToolCallsForResponsesLiteAPIKey(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
for _, passthrough := range []bool{false, true} {
|
||||
name := "managed"
|
||||
if passthrough {
|
||||
name = "passthrough"
|
||||
}
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
|
||||
c.Request.Header.Set("User-Agent", "codex_cli_rs/0.144.1")
|
||||
c.Request.Header.Set(responsesLiteHeader, "true")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_lite\",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" +
|
||||
"data: [DONE]\n\n",
|
||||
)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
|
||||
account := &Account{
|
||||
ID: 503, Name: "responses-lite-api-key", Platform: PlatformOpenAI, Type: AccountTypeAPIKey,
|
||||
Concurrency: 1, Status: StatusActive, Schedulable: true, RateMultiplier: f64p(1),
|
||||
Credentials: map[string]any{"api_key": "sk-test"},
|
||||
Extra: map[string]any{"openai_passthrough": passthrough},
|
||||
}
|
||||
body := []byte(`{
|
||||
"model":"gpt-5.6-terra","stream":true,"instructions":"test",
|
||||
"tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],
|
||||
"parallel_tool_calls":true,
|
||||
"input":[{"type":"message","role":"user","content":"hello"}]
|
||||
}`)
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, "true", upstream.lastReq.Header.Get(responsesLiteHeader))
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "tools").IsArray())
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "parallel_tool_calls").Exists())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "parallel_tool_calls").Bool())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -538,6 +538,34 @@ func TestOpenAIGatewayService_APIKeyStripsAllIndexedNamespacesBeforeFirstForward
|
||||
require.False(t, gjson.GetBytes(upstream.bodies[0], "input.1.namespace").Exists())
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceProactivelyStripsCrossProviderReasoningContent(t *testing.T) {
|
||||
body := []byte(`{"model":"gpt-5.5","stream":false,"store":true,"input":[` +
|
||||
`{"type":"message","role":"user","content":"one"},` +
|
||||
`{"type":"message","role":"assistant","content":"two"},` +
|
||||
`{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},` +
|
||||
`{"type":"function_call_output","call_id":"call_1","output":"ok"},` +
|
||||
`{"type":"message","role":"user","content":"five"},` +
|
||||
`{"type":"reasoning","summary":[{"type":"summary_text","text":"keep"}],"content":[{"type":"reasoning_text","text":"remove"}]}` +
|
||||
`]}`)
|
||||
upstream := &httpUpstreamRecorder{responses: []*http.Response{
|
||||
newOpenAIRejectedFieldTestResponse(http.StatusOK, `{"output":[],"usage":{"input_tokens":1,"output_tokens":1,"input_tokens_details":{"cached_tokens":0}}}`),
|
||||
}}
|
||||
|
||||
result, err := newOpenAIRejectedFieldTestService(upstream).Forward(
|
||||
context.Background(),
|
||||
newOpenAIRejectedFieldTestContext(body),
|
||||
newOpenAIRejectedFieldTestAccount(),
|
||||
body,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Len(t, upstream.bodies, 1, "reasoning content should be normalized before the first upstream request")
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(upstream.bodies[0], "input.5.type").String())
|
||||
require.False(t, gjson.GetBytes(upstream.bodies[0], "input.5.content").Exists())
|
||||
require.Equal(t, "keep", gjson.GetBytes(upstream.bodies[0], "input.5.summary.0.text").String())
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIHTTPStripsInputNamespacesBeforeFirstForward(t *testing.T) {
|
||||
accounts := []struct {
|
||||
name string
|
||||
|
||||
@@ -497,7 +497,7 @@ func TestOpenAIResponsesToolSchemaPlatformGate_APIKeyAndOAuth(t *testing.T) {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: accountType,
|
||||
})
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "object", gjson.GetBytes(normalized, "tools.0.parameters.type").String())
|
||||
@@ -508,7 +508,7 @@ func TestOpenAIResponsesToolSchemaPlatformGate_APIKeyAndOAuth(t *testing.T) {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeAPIKey,
|
||||
})
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, string(body), string(normalized))
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user