diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go index d1c389c7a..c463a870e 100644 --- a/backend/internal/service/antigravity_gateway_service_test.go +++ b/backend/internal/service/antigravity_gateway_service_test.go @@ -340,6 +340,65 @@ func TestAntigravityGatewayService_ForwardGemini_UsesConfiguredProjectFallback(t require.Equal(t, "configured-project", wrapped["project"]) } +func TestAntigravityGatewayService_ForwardGemini_ImageUsesDefaultMappingAndOAuth(t *testing.T) { + gin.SetMode(gin.TestMode) + writer := httptest.NewRecorder() + c, _ := gin.CreateTestContext(writer) + body := []byte(`{"contents":[{"role":"user","parts":[{"text":"draw a cat"}]}],"generationConfig":{"responseModalities":["TEXT","IMAGE"],"imageConfig":{"aspectRatio":"1:1"}}}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-3.1-flash-image:generateContent", bytes.NewReader(body)) + + upstream := &queuedHTTPUpstreamStub{ + responses: []*http.Response{{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"response\":{\"candidates\":[{\"content\":{\"parts\":[{\"inlineData\":{\"mimeType\":\"image/png\",\"data\":\"aGVsbG8=\"}}]},\"finishReason\":\"STOP\"}],\"usageMetadata\":{\"promptTokenCount\":1,\"candidatesTokenCount\":1}}}\n\n", + )), + }}, + onCall: func(req *http.Request, _ *queuedHTTPUpstreamStub) { + require.Equal(t, "Bearer test-access-token", req.Header.Get("Authorization")) + require.Equal(t, "application/json", req.Header.Get("Content-Type")) + require.Contains(t, req.URL.String(), "/v1internal:streamGenerateContent?alt=sse") + }, + } + svc := &AntigravityGatewayService{ + settingService: NewSettingService(&antigravitySettingRepoStub{}, &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}), + tokenProvider: &AntigravityTokenProvider{}, + httpUpstream: upstream, + } + account := &Account{ + ID: 104, + Name: "antigravity-image", + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "test-access-token", + "project_id": "test-project", + }, + } + + result, err := svc.ForwardGemini(context.Background(), c, account, "gemini-3.1-flash-image", "generateContent", true, body, false) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gemini-3.1-flash-image", result.Model) + require.Equal(t, "gemini-3.1-flash-image", result.UpstreamModel) + require.Equal(t, 1, result.ImageCount) + require.Len(t, upstream.requestBodies, 1) + + var wrapped map[string]any + require.NoError(t, json.Unmarshal(upstream.requestBodies[0], &wrapped)) + require.Equal(t, "test-project", wrapped["project"]) + require.Equal(t, "gemini-3.1-flash-image", wrapped["model"]) + request, ok := wrapped["request"].(map[string]any) + require.True(t, ok) + generationConfig, ok := request["generationConfig"].(map[string]any) + require.True(t, ok) + require.Equal(t, []any{"TEXT", "IMAGE"}, generationConfig["responseModalities"]) +} + func TestAntigravityGatewayService_ForwardGemini_PreservesServerSideToolInvocationConfig(t *testing.T) { gin.SetMode(gin.TestMode) body := []byte(`{"contents":[{"role":"user","parts":[{"text":"hello"}]}],"tools":[{"functionDeclarations":[{"name":"get_weather","parameters":{"type":"object","additionalProperties":false}}]},{"googleSearch":{}}],"toolConfig":{"includeServerSideToolInvocations":true}}`) diff --git a/backend/internal/service/openai_images_incomplete_test.go b/backend/internal/service/openai_images_incomplete_test.go index 6a8ba5622..ebd3c0e52 100644 --- a/backend/internal/service/openai_images_incomplete_test.go +++ b/backend/internal/service/openai_images_incomplete_test.go @@ -7,6 +7,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" "github.com/gin-gonic/gin" ) @@ -161,6 +162,91 @@ func TestImagesOAuthNonStreaming_ContentRefusalReturns400NoRetry(t *testing.T) { } } +func TestImagesOAuthNonStreaming_TextFallbackReturnsCapabilityError(t *testing.T) { + upstreamSSE := "event: response.created\n" + + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"r\",\"status\":\"in_progress\",\"model\":\"gpt-5.4-mini\",\"output\":[]}}\n\n" + + "event: response.output_text.delta\n" + + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"Here's a polished image prompt for your request.\"}\n\n" + + "event: response.completed\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r\",\"status\":\"completed\",\"model\":\"gpt-5.4-mini\",\"output\":[{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"Here's a polished image prompt for your request.\"}]}],\"tool_usage\":{\"image_gen\":{\"output_tokens\":0}}}}\n\n" + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} + + svc := &OpenAIGatewayService{} + _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + + var imgErr *OpenAIImagesUpstreamError + if !errors.As(err, &imgErr) { + t.Fatalf("expected *OpenAIImagesUpstreamError, got %T: %v", err, err) + } + if imgErr.StatusCode != http.StatusBadGateway { + t.Fatalf("text fallback should be retryable 502, got %d", imgErr.StatusCode) + } + if imgErr.Code != "image_generation_unavailable" { + t.Fatalf("text fallback should identify missing image execution, got %q", imgErr.Code) + } +} + +func TestImagesOAuthStreaming_TextFallbackReturnsCapabilityError(t *testing.T) { + upstreamSSE := "event: response.output_text.delta\n" + + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"Here's a polished image prompt for your request.\"}\n\n" + + "event: response.completed\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r\",\"status\":\"completed\",\"model\":\"gpt-5.4-mini\",\"output\":[]}}\n\n" + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} + + svc := &OpenAIGatewayService{} + _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2") + + var imgErr *OpenAIImagesUpstreamError + if !errors.As(err, &imgErr) { + t.Fatalf("expected *OpenAIImagesUpstreamError, got %T: %v", err, err) + } + if imgErr.StatusCode != http.StatusBadGateway { + t.Fatalf("streaming text fallback should be retryable 502, got %d", imgErr.StatusCode) + } + if imgErr.Code != "image_generation_unavailable" { + t.Fatalf("streaming text fallback should identify missing image execution, got %q", imgErr.Code) + } + if strings.Contains(rec.Body.String(), "event: error") { + t.Fatal("retryable text fallback must remain unflushed for failover") + } +} + +func TestImagesOAuthStreaming_SplitSafetyRefusalReturns400(t *testing.T) { + upstreamSSE := "event: response.output_text.delta\n" + + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"安全系\"}\n\n" + + "event: response.output_text.delta\n" + + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"统拒绝生成\"}\n\n" + + "event: response.completed\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r\",\"status\":\"completed\",\"output\":[]}}\n\n" + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} + + svc := &OpenAIGatewayService{} + _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2") + + var imgErr *OpenAIImagesUpstreamError + if !errors.As(err, &imgErr) { + t.Fatalf("expected *OpenAIImagesUpstreamError, got %T: %v", err, err) + } + if imgErr.StatusCode != http.StatusBadRequest || imgErr.Code != "content_policy_violation" { + t.Fatalf("split safety refusal should remain a content-policy 400, got status=%d code=%q", imgErr.StatusCode, imgErr.Code) + } + if !strings.Contains(rec.Body.String(), "event: error") { + t.Fatal("content-policy refusal must reach the streaming client") + } +} + // extractOpenAIImagesModelRefusal:真空响应(无文字)返回空串。 func TestExtractModelRefusal_EmptyWhenNoText(t *testing.T) { body := "data: {\"type\":\"response.completed\",\"response\":{\"output\":[],\"tool_usage\":{\"image_gen\":{\"output_tokens\":0}}}}\n\n" diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 79680e66e..7e32dde8d 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -609,13 +609,10 @@ func openAIImagesUpstreamErrorFromSSEPayload(payload []byte) *OpenAIImagesUpstre } } -// extractOpenAIImagesModelRefusal 从上游 SSE 响应体提取「模型未出图、改用文字拒绝」 -// 的拒绝文本(内容审核场景)。 -// -// 上游 response.completed 无图时,模型常以 output_text / message 形式输出拒绝说明 -// (如“被安全系统判定为不适合生成”)。这类失败是内容策略拦截,重试/换账号均无效, -// 应把该文本作为内容策略错误透传给客户端。返回空串表示无文字输出(真空响应)。 -func extractOpenAIImagesModelRefusal(body []byte) string { +// extractOpenAIImagesModelText collects textual terminal output from an image +// request. A text response is evidence that the image tool did not produce an +// image; its semantics still need classification before choosing a client error. +func extractOpenAIImagesModelText(body []byte) string { var b strings.Builder collect := func(s string) { if s = strings.TrimSpace(s); s != "" { @@ -625,16 +622,14 @@ func extractOpenAIImagesModelRefusal(body []byte) string { _, _ = b.WriteString(s) } } - forEachOpenAISSEDataPayload(string(body), func(payload []byte) { + consumePayload := func(payload []byte) { if !gjson.ValidBytes(payload) { return } switch gjson.GetBytes(payload, "type").String() { case "response.output_text.delta": - // 流式文本增量。 collect(gjson.GetBytes(payload, "delta").String()) case "response.completed", "response.output_item.done": - // 终态里的 message/output_text。 gjson.GetBytes(payload, "response.output").ForEach(func(_, item gjson.Result) bool { if item.Get("type").String() == "message" { item.Get("content").ForEach(func(_, part gjson.Result) bool { @@ -655,14 +650,67 @@ func extractOpenAIImagesModelRefusal(body []byte) string { }) } } - }) - refusal := strings.TrimSpace(b.String()) - // 截断过长文本,避免把整段模型输出塞进错误响应。 - const maxRefusal = 600 - if len(refusal) > maxRefusal { - refusal = refusal[:maxRefusal] } - return refusal + if gjson.ValidBytes(body) { + consumePayload(body) + } else { + forEachOpenAISSEDataPayload(string(body), consumePayload) + } + text := strings.TrimSpace(b.String()) + const maxText = 600 + if len(text) > maxText { + return text[:maxText] + } + return text +} + +func isOpenAIImagesContentPolicyRefusal(text string) bool { + lower := strings.ToLower(text) + for _, marker := range []string{ + "content policy", "content_policy", "content filter", "content_filter", + "safety system", "safety policy", "safety violation", "moderation", + "安全系统", "安全策略", "安全政策", "内容政策", "内容审核", "违规内容", "不适合生成", + } { + if strings.Contains(lower, marker) { + return true + } + } + return false +} + +// extractOpenAIImagesModelRefusal returns only text with an explicit safety or +// moderation signal. Plain prompt suggestions are capability failures instead. +func extractOpenAIImagesModelRefusal(body []byte) string { + text := extractOpenAIImagesModelText(body) + if !isOpenAIImagesContentPolicyRefusal(text) { + return "" + } + return text +} + +func openAIImagesTextFallbackError(body []byte) *OpenAIImagesUpstreamError { + return openAIImagesTextFallbackErrorForText(extractOpenAIImagesModelText(body)) +} + +func openAIImagesTextFallbackErrorForText(text string) *OpenAIImagesUpstreamError { + text = strings.TrimSpace(text) + if text == "" { + return nil + } + if isOpenAIImagesContentPolicyRefusal(text) { + return &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadRequest, + ErrorType: "image_generation_user_error", + Code: "content_policy_violation", + Message: sanitizeUpstreamErrorMessage(text), + } + } + return &OpenAIImagesUpstreamError{ + StatusCode: http.StatusBadGateway, + ErrorType: "upstream_error", + Code: "image_generation_unavailable", + Message: "Upstream did not execute image generation", + } } // summarizeOpenAIImagesNoOutputBody 从上游 SSE 响应体提取诊断摘要,用于软失败时 @@ -1267,26 +1315,14 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse( } return OpenAIUsage{}, 0, nil, upstreamErr } - // 软失败兜底:上游无图。先区分两种情形(实测真因,见下): - // - // (A) 内容审核拒绝:模型未出图,但输出了文字拒绝(response.completed 里带 - // output_text / message,内容如“被安全系统判定为不适合生成”)。这是用户 - // prompt 触发 OpenAI 内容策略,模型主动拒绝改用文字回应。**换账号/重试均无效** - // (内容层拦截,与账号/承载模型无关),应把拒绝理由作为 400 透传给客户端, - // 避免无谓地重试 + 消耗其它账号配额,且让客户端拿到可读的拒绝原因。 - // (B) 真空响应:既无图也无任何文字输出(罕见,如偶发路由到 gpt-5.x-mini、 - // image_gen 工具未执行)。这是上游的概率性失败,此时才按可重试处理。 - if refusal := extractOpenAIImagesModelRefusal(body); refusal != "" { - refusalErr := &OpenAIImagesUpstreamError{ - StatusCode: http.StatusBadRequest, - ErrorType: "image_generation_user_error", - Code: "content_policy_violation", - Message: sanitizeUpstreamErrorMessage(refusal), + if textFallbackErr := openAIImagesTextFallbackError(body); textFallbackErr != nil { + setOpsUpstreamError(c, textFallbackErr.clientStatusCode(), textFallbackErr.clientMessage(), summarizeOpenAIImagesNoOutputBody(body)) + if !IsOpenAIImagesRetryableUpstreamError(textFallbackErr) { + writeOpenAIImagesUpstreamErrorResponse(c, textFallbackErr) } - setOpsUpstreamError(c, http.StatusBadRequest, refusalErr.clientMessage(), summarizeOpenAIImagesNoOutputBody(body)) - writeOpenAIImagesUpstreamErrorResponse(c, refusalErr) - return OpenAIUsage{}, 0, nil, refusalErr + return OpenAIUsage{}, 0, nil, textFallbackErr } + // 真空响应:既无图也无文字输出。它保持短暂可重试语义,优先同账号重试。 // (B) 真空响应:记录上游诊断摘要到 ops(last_event/status/model/body 片段)便于 // 排查,并返回 UpstreamFailoverError 触发重试。因实测为「同账号概率性失败」,优先 // RetryableOnSameAccount 同账号快速重试(默认 3 次,大概率某次正常出图),用尽后 @@ -1344,6 +1380,17 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( pendingResults := make([]openAIResponsesImageResult, 0, 1) pendingSeen := make(map[string]struct{}) streamMeta := openAIResponsesImageResult{Model: strings.TrimSpace(fallbackModel)} + var fallbackText strings.Builder + appendFallbackText := func(text string) { + if text == "" || fallbackText.Len() >= 600 { + return + } + remaining := 600 - fallbackText.Len() + if len(text) > remaining { + text = text[:remaining] + } + _, _ = fallbackText.WriteString(text) + } var createdAt int64 clientDisconnected := false lastDownstreamWriteAt := time.Now() @@ -1370,6 +1417,9 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( createdAt = eventCreatedAt } } + if gjson.GetBytes(dataBytes, "type").String() == "response.output_text.delta" { + appendFallbackText(gjson.GetBytes(dataBytes, "delta").String()) + } switch gjson.GetBytes(dataBytes, "type").String() { case "response.image_generation_call.partial_image": b64 := strings.TrimSpace(gjson.GetBytes(dataBytes, "partial_image_b64").String()) @@ -1434,9 +1484,21 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( } reconcileOpenAIResponsesImageResultSizes(finalResults, nil) if len(finalResults) == 0 { + textFallbackErr := openAIImagesTextFallbackErrorForText(fallbackText.String()) + if textFallbackErr == nil { + textFallbackErr = openAIImagesTextFallbackError(dataBytes) + } + if textFallbackErr != nil { + retryable := IsOpenAIImagesRetryableUpstreamError(textFallbackErr) + setOpsUpstreamError(c, textFallbackErr.clientStatusCode(), textFallbackErr.clientMessage(), summarizeOpenAIImagesNoOutputBody(dataBytes)) + if !retryable && !clientDisconnected { + s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, "error", buildOpenAIImagesStreamErrorBodyFromUpstream(textFallbackErr)) + } + processDataErr = textFallbackErr + processDataDone = true + return + } outputErr := fmt.Errorf("upstream did not return image output") - // 软失败:response.completed 事件里没有图片。记录上游诊断摘要到 ops, - // 与非流式路径保持一致,避免上游响应信息丢失。 setOpsUpstreamError(c, http.StatusBadGateway, "upstream did not return image output", summarizeOpenAIImagesNoOutputBody(dataBytes)) s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, "error", buildOpenAIImagesStreamErrorBody(outputErr.Error())) processDataErr = outputErr @@ -1718,6 +1780,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( } upstreamReq.Header.Set("Content-Type", "application/json") upstreamReq.Header.Set("Accept", "text/event-stream") + upstreamReq.Header.Set("OpenAI-Beta", "responses=experimental") proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -1853,6 +1916,25 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( }, nil } +const ( + openAIImagesOAuthUnavailableCooldown = 30 * time.Minute + openAIImagesOAuthUnavailableReason = "openai_images_oauth_tool_unavailable" +) + +func (s *OpenAIGatewayService) coolOpenAIImagesOAuthTool(ctx context.Context, account *Account) { + if s == nil || s.accountRepo == nil || account == nil || account.Platform != PlatformOpenAI { + return + } + stateCtx, cancel := openAIAccountStateContext(ctx) + defer cancel() + resetAt := time.Now().Add(openAIImagesOAuthUnavailableCooldown) + if err := s.accountRepo.SetModelRateLimit(stateCtx, account.ID, openAIImageGenerationRateLimitKey, resetAt, openAIImagesOAuthUnavailableReason); err != nil { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Images OAuth tool cooldown write failed account_id=%d error=%v", account.ID, err) + return + } + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Images OAuth tool unavailable account_id=%d reset_in=%s", account.ID, time.Until(resetAt).Truncate(time.Second)) +} + func (s *OpenAIGatewayService) handleOpenAIImagesOAuthResponseError( ctx context.Context, c *gin.Context, @@ -1932,11 +2014,25 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthResponseError( Message: upstreamErr.clientMessage(), }) + responseBody := openAIImagesUpstreamErrorResponseBody(upstreamErr) + if upstreamErr.Code == "image_generation_unavailable" { + s.coolOpenAIImagesOAuthTool(ctx, account) + if responseWritten { + return err + } + return s.newOpenAIAccountFailoverError( + account, + upstreamErr.StatusCode, + headers, + responseBody, + upstreamErr.clientMessage(), + false, + false, + ) + } if !retryable || responseWritten { return err } - - responseBody := openAIImagesUpstreamErrorResponseBody(upstreamErr) shouldDisable := s.handleOpenAIAccountUpstreamError(ctx, account, upstreamErr.StatusCode, headers, responseBody, requestedModel) return s.newOpenAIAccountFailoverError( account, diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index 99afeba6a..6761c5513 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -847,7 +847,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthPassesNAndReturnsAllImages(t *te require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) require.Equal(t, "acct-123", upstream.lastReq.Header.Get("chatgpt-account-id")) - require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) + require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) @@ -1963,6 +1963,21 @@ func TestBuildOpenAIImagesResponsesRequest_PassesThroughNForMultiImageModels(t * require.Equal(t, "draw a cat", gjson.GetBytes(body, "input.0.content.0.text").String()) } +func TestBuildOpenAIImagesResponsesRequest_ForcesImageToolChoice(t *testing.T) { + parsed := &OpenAIImagesRequest{ + Endpoint: openAIImagesGenerationsEndpoint, + Model: "gpt-image-2", + Prompt: "draw a cat", + } + + body, err := buildOpenAIImagesResponsesRequest(parsed, "gpt-image-2") + require.NoError(t, err) + require.NotNil(t, body) + require.Equal(t, "image_generation", gjson.GetBytes(body, "tool_choice.type").String()) + require.Equal(t, "image_generation", gjson.GetBytes(body, "tools.0.type").String()) + require.Equal(t, "gpt-image-2", gjson.GetBytes(body, "tools.0.model").String()) +} + func TestBuildOpenAIImagesResponsesRequest_DoesNotPassNForDallE3(t *testing.T) { parsed := &OpenAIImagesRequest{ Endpoint: openAIImagesGenerationsEndpoint, diff --git a/backend/internal/service/ratelimit_service_openai_image_test.go b/backend/internal/service/ratelimit_service_openai_image_test.go index 76d3bb764..26714cf12 100644 --- a/backend/internal/service/ratelimit_service_openai_image_test.go +++ b/backend/internal/service/ratelimit_service_openai_image_test.go @@ -119,3 +119,53 @@ func TestOpenAIGatewayServiceForwardImages_ImageRateLimitReturnsFailoverAndCools require.Len(t, repo.modelRateLimitCalls, 1) require.Equal(t, openAIImageGenerationRateLimitKey, repo.modelRateLimitCalls[0].scope) } + +func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *testing.T) { + gin.SetMode(gin.TestMode) + repo := &modelNotFoundAccountRepoStub{} + body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`) + upstreamSSE := "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r\",\"status\":\"completed\",\"model\":\"gpt-5.4-mini\",\"output\":[{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"Here's a polished image prompt for your request.\"}]}]}}\n\n" + + req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = req + + svc := &OpenAIGatewayService{ + accountRepo: repo, + httpUpstream: &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }, + }, + } + parsed, err := svc.ParseOpenAIImagesRequest(c, body) + require.NoError(t, err) + account := &Account{ + ID: 205, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token-123", + }, + } + + before := time.Now() + result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") + + require.Nil(t, result) + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.False(t, failoverErr.RetryableOnSameAccount) + require.Len(t, repo.modelRateLimitCalls, 1) + call := repo.modelRateLimitCalls[0] + require.Equal(t, account.ID, call.accountID) + require.Equal(t, openAIImageGenerationRateLimitKey, call.scope) + require.Equal(t, openAIImagesOAuthUnavailableReason, call.reason) + require.WithinDuration(t, before.Add(openAIImagesOAuthUnavailableCooldown), call.resetAt, time.Second) +}