diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go index b2df1435..f6cd1404 100644 --- a/backend/internal/service/gateway_service_benchmark_test.go +++ b/backend/internal/service/gateway_service_benchmark_test.go @@ -124,6 +124,64 @@ func BenchmarkOpenAIResponses_LargeInputDecodeMap(b *testing.B) { } } +func BenchmarkOpenAIResponses_LargeInputRawPatch(b *testing.B) { + for _, size := range benchmarkBodySizes() { + b.Run(size.name, func(b *testing.B) { + body := buildLargeOpenAIResponsesBody(size.bytes) + + b.SetBytes(int64(len(body))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + view := newOpenAIRequestView(body) + view.MarkPatchSet("instructions", "You are a helpful coding assistant.") + view.MarkPatchSet("reasoning.effort", "none") + patched, err := view.ApplyPatches() + if err != nil { + b.Fatalf("应用 OpenAI raw patch 失败: %v", err) + } + benchmarkIntSink = len(patched) + } + }) + } +} + +func BenchmarkOpenAIResponses_LargeInputImageBillingRaw(b *testing.B) { + for _, size := range benchmarkBodySizes() { + b.Run(size.name, func(b *testing.B) { + body := buildLargeOpenAIResponsesImageToolBody(size.bytes) + + b.SetBytes(int64(len(body))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, "gpt-5.4") + if err != nil { + b.Fatalf("解析 OpenAI 图片计费配置失败: %v", err) + } + benchmarkStringSink = cfg.Model + cfg.SizeTier + cfg.InputSize + } + }) + } +} + +func BenchmarkOpenAIResponses_LargeInputEmptyBase64Guard(b *testing.B) { + for _, size := range benchmarkBodySizes() { + b.Run(size.name, func(b *testing.B) { + body := buildLargeOpenAIResponsesBody(size.bytes) + + b.SetBytes(int64(len(body))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if openAIRequestBodyMayContainEmptyBase64InputImage(body) { + benchmarkIntSink++ + } + } + }) + } +} + func BenchmarkOpenAIResponses_LargeInputFunctionCallValidation(b *testing.B) { for _, size := range benchmarkBodySizes() { b.Run(size.name, func(b *testing.B) { @@ -264,3 +322,20 @@ func buildLargeOpenAIResponsesToolContinuationBody(targetBytes int) []byte { builder.WriteString(`]}`) return []byte(builder.String()) } + +func buildLargeOpenAIResponsesImageToolBody(targetBytes int) []byte { + var builder strings.Builder + builder.Grow(targetBytes + 1024) + builder.WriteString(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[`) + for i := 0; builder.Len() < targetBytes; i++ { + if i > 0 { + builder.WriteByte(',') + } + builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`) + builder.WriteString(strings.Repeat("openai image billing payload ", 48)) + builder.WriteString(strconv.Itoa(i)) + builder.WriteString(`"}]}`) + } + builder.WriteString(`]}`) + return []byte(builder.String()) +} diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go index 4aca1239..80b6c66d 100644 --- a/backend/internal/service/image_generation_intent.go +++ b/backend/internal/service/image_generation_intent.go @@ -1,7 +1,6 @@ package service import ( - "encoding/json" "strings" "github.com/tidwall/gjson" @@ -91,7 +90,7 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { } found := false tools.ForEach(func(_, item gjson.Result) bool { - if strings.TrimSpace(item.Get("type").String()) == "image_generation" { + if openAIJSONString(item.Get("type")) == "image_generation" { found = true return false } @@ -100,6 +99,36 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { return found } +func openAIRequestBodyHasImageGenerationTool(body []byte) bool { + if len(body) == 0 || !gjson.ValidBytes(body) { + return false + } + return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools")) +} + +func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool { + if len(body) == 0 || !gjson.ValidBytes(body) { + return false + } + tools := gjson.GetBytes(body, "tools") + if !tools.IsArray() { + return false + } + needsNormalization := false + tools.ForEach(func(_, item gjson.Result) bool { + if openAIJSONString(item.Get("type")) != "image_generation" { + return true + } + // 只有旧字段需要迁移时才进入 map 修改,纯计费读取保持 raw 路径。 + if item.Get("format").Exists() || item.Get("compression").Exists() { + needsNormalization = true + return false + } + return true + }) + return needsNormalization +} + func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool { if !choice.Exists() { return false @@ -159,17 +188,6 @@ func apiKeyGroup(apiKey *APIKey) *Group { return apiKey.Group } -func cloneRequestMapForImageIntent(body []byte) map[string]any { - if len(body) == 0 { - return nil - } - var out map[string]any - if err := json.Unmarshal(body, &out); err != nil { - return nil - } - return out -} - type OpenAIResponsesImageBillingConfig struct { Model string SizeTier string @@ -225,8 +243,43 @@ func resolveOpenAIResponsesImageBillingConfigFromBody(body []byte, fallbackModel } func resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body []byte, fallbackModel string) (OpenAIResponsesImageBillingConfig, error) { - reqBody := cloneRequestMapForImageIntent(body) - return resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, fallbackModel) + imageModel := "" + imageSize := "" + hasImageTool := false + if len(body) > 0 && gjson.ValidBytes(body) { + tools := gjson.GetBytes(body, "tools") + if tools.IsArray() { + tools.ForEach(func(_, item gjson.Result) bool { + if openAIJSONString(item.Get("type")) != "image_generation" { + return true + } + hasImageTool = true + imageModel = openAIJSONString(item.Get("model")) + imageSize = openAIJSONString(item.Get("size")) + return false + }) + } + if imageSize == "" { + imageSize = openAIJSONString(gjson.GetBytes(body, "size")) + } + if imageModel == "" { + bodyModel := openAIJSONString(gjson.GetBytes(body, "model")) + if isOpenAIImageBillingModelAlias(bodyModel) || !hasImageTool { + imageModel = bodyModel + } + } + } + if imageModel == "" && hasImageTool { + imageModel = "gpt-image-2" + } + if imageModel == "" { + imageModel = strings.TrimSpace(fallbackModel) + } + return OpenAIResponsesImageBillingConfig{ + Model: imageModel, + SizeTier: normalizeOpenAIImageSizeTier(imageSize), + InputSize: imageSize, + }, nil } func isOpenAIImageBillingModelAlias(model string) bool { @@ -236,3 +289,10 @@ func isOpenAIImageBillingModelAlias(model string) bool { } return isOpenAIImageGenerationModel(normalized) || strings.Contains(normalized, "image") } + +func openAIJSONString(value gjson.Result) string { + if value.Type != gjson.String { + return "" + } + return strings.TrimSpace(value.String()) +} diff --git a/backend/internal/service/image_generation_intent_test.go b/backend/internal/service/image_generation_intent_test.go index 4621e9d9..59aab39c 100644 --- a/backend/internal/service/image_generation_intent_test.go +++ b/backend/internal/service/image_generation_intent_test.go @@ -84,6 +84,17 @@ func TestResolveOpenAIResponsesImageBillingConfigToolModelWins(t *testing.T) { require.Equal(t, "2K", imageSize) } +func TestResolveOpenAIResponsesImageBillingConfigFromBodyIgnoresUnrelatedLargeInput(t *testing.T) { + cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody( + []byte(`{"model":"mapped-text-model","tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":1e1000000}]}]}`), + "requested-model", + ) + require.NoError(t, err) + require.Equal(t, "gpt-image-2", cfg.Model) + require.Equal(t, "2K", cfg.SizeTier) + require.Equal(t, "2048x1152", cfg.InputSize) +} + func TestResolveOpenAIResponsesImageBillingConfigSupportsOfficialAndCustomSizes(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 87ffb2cb..d17b85d0 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -2397,172 +2397,83 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco return s.forwardOpenAIPassthrough(ctx, c, account, originalBody, reqModel, reasoningEffort, reqStream, startTime) } - reqBody, err := requestView.Decode(c) - if err != nil { - return nil, err + bodyModified := false + var reqBody map[string]any + ensureReqBody := func() (map[string]any, error) { + if requestView.HasPatches() { + patchedBody, patchErr := requestView.ApplyPatches() + if patchErr != nil { + return nil, patchErr + } + body = patchedBody + requestView = newOpenAIRequestView(body) + reqBody = nil + bodyModified = false + } + if reqBody != nil { + return reqBody, nil + } + decoded, decodeErr := requestView.Decode(c) + if decodeErr != nil { + return nil, decodeErr + } + reqBody = decoded + return reqBody, nil + } + markPatchSet := func(path string, value any) { + bodyModified = true + if requestView.patchesDisabled { + if reqBody != nil { + setOpenAIRequestMapPath(reqBody, path, value) + } + return + } + requestView.MarkPatchSet(path, value) + } + markPatchDelete := func(path string) { + bodyModified = true + if requestView.patchesDisabled { + if reqBody != nil { + deleteOpenAIRequestMapPath(reqBody, path) + } + return + } + requestView.MarkPatchDelete(path) + } + disablePatch := func() { + requestView.DisablePatches() + } + markDecodedModified := func() { + bodyModified = true + disablePatch() } - if v, ok := reqBody["model"].(string); ok { - reqModel = v - originalModel = reqModel - } - if v, ok := reqBody["stream"].(bool); ok { - reqStream = v - } - if promptCacheKey == "" { - if v, ok := reqBody["prompt_cache_key"].(string); ok { - promptCacheKey = strings.TrimSpace(v) - } - } apiKey := getAPIKeyFromContext(c) imageGenerationAllowed := GroupAllowsImageGeneration(nil) if apiKey != nil { imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group) } codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) - if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed { + imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) + if imageIntent && !imageGenerationAllowed { MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) - c.JSON(http.StatusForbidden, gin.H{ - "error": gin.H{ - "type": "permission_error", - "message": ImageGenerationPermissionMessage(), - }, - }) + c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}}) return nil, errors.New("image generation disabled for group") } - // Track if body needs re-serialization - bodyModified := false - // 单字段补丁快速路径:只要整个变更集最终可归约为同一路径的 set/delete,就避免全量 Marshal。 - patchDisabled := false - patchHasOp := false - patchDelete := false - patchPath := "" - var patchValue any - markPatchSet := func(path string, value any) { - if strings.TrimSpace(path) == "" { - patchDisabled = true - return - } - if patchDisabled { - return - } - if !patchHasOp { - patchHasOp = true - patchDelete = false - patchPath = path - patchValue = value - return - } - if patchDelete || patchPath != path { - patchDisabled = true - return - } - patchValue = value - } - markPatchDelete := func(path string) { - if strings.TrimSpace(path) == "" { - patchDisabled = true - return - } - if patchDisabled { - return - } - if !patchHasOp { - patchHasOp = true - patchDelete = true - patchPath = path - return - } - if !patchDelete || patchPath != path { - patchDisabled = true - } - } - disablePatch := func() { - patchDisabled = true - } - - // 非透传模式下,instructions 为空时注入默认指令。 - if isInstructionsEmpty(reqBody) && !compatMessagesBridge { - reqBody["instructions"] = "You are a helpful coding assistant." - bodyModified = true + instructions := gjson.GetBytes(body, "instructions") + instructionsEmpty := !instructions.Exists() || instructions.Type != gjson.String || strings.TrimSpace(instructions.String()) == "" + if instructionsEmpty && !compatMessagesBridge { markPatchSet("instructions", "You are a helpful coding assistant.") } - if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(reqBody) { - bodyModified = true - disablePatch() - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client") - } - - if normalizeOpenAIResponsesImageGenerationTools(reqBody) { - bodyModified = true - disablePatch() - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload") - } - if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(reqBody) { - bodyModified = true - disablePatch() - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions") - } - - // 对所有请求执行模型映射(包含 Codex CLI)。 billingModel := account.GetMappedModel(reqModel) if billingModel != reqModel { logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", reqModel, billingModel, account.Name, isCodexCLI) - reqBody["model"] = billingModel - bodyModified = true + reqModel = billingModel markPatchSet("model", billingModel) } upstreamModel := billingModel - if imageGenerationAllowed && normalizeOpenAIResponsesImageOnlyModel(reqBody) { - bodyModified = true - disablePatch() - if model, ok := reqBody["model"].(string); ok { - upstreamModel = strings.TrimSpace(model) - } - logger.LegacyPrintf( - "service.openai_gateway", - "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s", - reqModel, - billingModel, - upstreamModel, - ) - } - if err := validateOpenAIResponsesImageModel(reqBody, upstreamModel); err != nil { - setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "") - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": err.Error(), - "param": "model", - }, - }) - return nil, err - } - if hasOpenAIImageGenerationTool(reqBody) { - logger.LegacyPrintf( - "service.openai_gateway", - "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", - reqModel, - upstreamModel, - account.Type, - ) - } - if err := validateCodexSparkInput(reqBody, upstreamModel); err != nil { - setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "") - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": err.Error(), - "param": "input", - }, - }) - return nil, err - } - - // Compact-only model 映射:仅在 /responses/compact 路径生效,且优先级高于 - // OAuth 模型规范化(避免 OAuth 规范化覆盖 compact-only 自定义模型)。 isCompactRequest := isOpenAIResponsesCompactPath(c) compactMapped := false if isCompactRequest { @@ -2570,69 +2481,99 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if compactMappedModel != "" && compactMappedModel != billingModel { compactMapped = true upstreamModel = compactMappedModel - reqBody["model"] = compactMappedModel - bodyModified = true + reqModel = compactMappedModel markPatchSet("model", compactMappedModel) logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Compact model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", billingModel, compactMappedModel, account.Name, isCodexCLI) } } - - // OpenAI OAuth 账号走 ChatGPT internal Codex endpoint,需要将模型名规范化为 - // 上游可识别的 Codex/GPT 系列。API Key 账号则应保留原始/映射后的模型名, - // 以兼容自定义 base_url 的 OpenAI-compatible 上游。 - if model, ok := reqBody["model"].(string); ok { - if !compactMapped { - upstreamModel = normalizeOpenAIModelForUpstream(account, model) - if upstreamModel != "" && upstreamModel != model { - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", - model, upstreamModel, account.Name, account.Type, isCodexCLI) - reqBody["model"] = upstreamModel - bodyModified = true - markPatchSet("model", upstreamModel) - } + if !compactMapped { + modelForNormalize := reqModel + if modelForNormalize == "" { + modelForNormalize = requestView.Model } - - // 移除 gpt-5.2-codex 以下的版本 verbosity 参数 - // 确保高版本模型向低版本模型映射不报错 - if !SupportsVerbosity(upstreamModel) { - if text, ok := reqBody["text"].(map[string]any); ok { - if _, exists := text["verbosity"]; exists { - delete(text, "verbosity") - bodyModified = true - markPatchDelete("text.verbosity") - } - } + upstreamModel = normalizeOpenAIModelForUpstream(account, modelForNormalize) + if upstreamModel != "" && upstreamModel != modelForNormalize { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", modelForNormalize, upstreamModel, account.Name, account.Type, isCodexCLI) + reqModel = upstreamModel + markPatchSet("model", upstreamModel) } } + if strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()) == "minimal" { + markPatchSet("reasoning.effort", "none") + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name) + } - // 规范化 reasoning.effort 参数(minimal -> none),与上游允许值对齐。 - if reasoning, ok := reqBody["reasoning"].(map[string]any); ok { - if effort, ok := reasoning["effort"].(string); ok && effort == "minimal" { - reasoning["effort"] = "none" - bodyModified = true - markPatchSet("reasoning.effort", "none") - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name) + imageIntent = imageIntent || IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, nil) || isOpenAIImageGenerationModel(upstreamModel) + if imageIntent && !imageGenerationAllowed { + MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}}) + return nil, errors.New("image generation disabled for group") + } + + if imageGenerationAllowed && (codexImageGenerationBridgeEnabled || isOpenAIImageGenerationModel(requestView.Model) || openAIRequestBodyImageGenerationToolNeedsNormalization(body) || isOpenAIImageGenerationModel(upstreamModel)) { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client") + } + if normalizeOpenAIResponsesImageGenerationTools(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload") + } + if normalizeOpenAIResponsesImageOnlyModel(decoded) { + markDecodedModified() + if model, ok := decoded["model"].(string); ok { + upstreamModel = strings.TrimSpace(model) + } + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s", requestView.Model, billingModel, upstreamModel) + } + if err := validateOpenAIResponsesImageModel(decoded, upstreamModel); err != nil { + setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "") + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "model"}}) + return nil, err + } + if hasOpenAIImageGenerationTool(decoded) { + imageIntent = true + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type) + } + if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions") + } + } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type) + } + + if isCodexSparkModel(upstreamModel) && openAIRequestBodyMayContainImageInput(body) { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if err := validateCodexSparkInput(decoded, upstreamModel); err != nil { + setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "") + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "input"}}) + return nil, err } } if account.Type == AccountTypeOAuth { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } codexResult := codexTransformResult{} if compatMessagesBridge { - codexResult = applyCodexOAuthTransformWithOptions(reqBody, codexOAuthTransformOptions{ - IsCodexCLI: isCodexCLI, - IsCompact: isCompactRequest, - SkipDefaultInstructions: true, - PreserveToolCallIDs: true, - }) - ensureCodexOAuthInstructionsField(reqBody) - bodyModified = true - disablePatch() + codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{IsCodexCLI: isCodexCLI, IsCompact: isCompactRequest, SkipDefaultInstructions: true, PreserveToolCallIDs: true}) + ensureCodexOAuthInstructionsField(decoded) + markDecodedModified() } else { - codexResult = applyCodexOAuthTransform(reqBody, isCodexCLI, isCompactRequest) + codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest) } if codexResult.Modified { - bodyModified = true - disablePatch() + markDecodedModified() } if codexResult.NormalizedModel != "" { upstreamModel = codexResult.NormalizedModel @@ -2642,90 +2583,57 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } } - // Handle max_output_tokens based on platform and account type + if !SupportsVerbosity(upstreamModel) && gjson.GetBytes(body, "text.verbosity").Exists() { + markPatchDelete("text.verbosity") + } + if !isCodexCLI { - if maxOutputTokens, hasMaxOutputTokens := reqBody["max_output_tokens"]; hasMaxOutputTokens { + maxOutputTokens := gjson.GetBytes(body, "max_output_tokens") + if maxOutputTokens.Exists() { switch account.Platform { case PlatformOpenAI: - // For OpenAI API Key, remove max_output_tokens (not supported) - // For OpenAI OAuth (Responses API), keep it (supported) if account.Type == AccountTypeAPIKey { - delete(reqBody, "max_output_tokens") - bodyModified = true markPatchDelete("max_output_tokens") } case PlatformAnthropic: - // For Anthropic (Claude), convert to max_tokens - delete(reqBody, "max_output_tokens") - markPatchDelete("max_output_tokens") - if _, hasMaxTokens := reqBody["max_tokens"]; !hasMaxTokens { - reqBody["max_tokens"] = maxOutputTokens - disablePatch() + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr } - bodyModified = true + delete(decoded, "max_output_tokens") + if _, hasMaxTokens := decoded["max_tokens"]; !hasMaxTokens { + decoded["max_tokens"] = maxOutputTokens.Value() + } + markDecodedModified() case PlatformGemini: - // For Gemini, remove (will be handled by Gemini-specific transform) - delete(reqBody, "max_output_tokens") - bodyModified = true markPatchDelete("max_output_tokens") default: - // For unknown platforms, remove to be safe - delete(reqBody, "max_output_tokens") - bodyModified = true markPatchDelete("max_output_tokens") } } - - // Also handle max_completion_tokens (similar logic) - if _, hasMaxCompletionTokens := reqBody["max_completion_tokens"]; hasMaxCompletionTokens { - if account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI { - delete(reqBody, "max_completion_tokens") - bodyModified = true - markPatchDelete("max_completion_tokens") - } + if gjson.GetBytes(body, "max_completion_tokens").Exists() && (account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI) { + markPatchDelete("max_completion_tokens") } - - // Remove unsupported fields (not supported by upstream OpenAI API) - unsupportedFields := []string{"prompt_cache_retention", "safety_identifier"} - for _, unsupportedField := range unsupportedFields { - if _, has := reqBody[unsupportedField]; has { - delete(reqBody, unsupportedField) - bodyModified = true + for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} { + if gjson.GetBytes(body, unsupportedField).Exists() { markPatchDelete(unsupportedField) } } } - - // 仅在 WSv2 模式保留 previous_response_id,其他模式(HTTP/WSv1)统一过滤。 - // 注意:该规则同样适用于 Codex CLI 请求,避免 WSv1 向上游透传不支持字段。 - if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 { - if _, has := reqBody["previous_response_id"]; has { - delete(reqBody, "previous_response_id") - bodyModified = true - markPatchDelete("previous_response_id") + if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 && gjson.GetBytes(body, "previous_response_id").Exists() { + markPatchDelete("previous_response_id") + } + if openAIRequestBodyMayContainEmptyBase64InputImage(body) { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(decoded) { + markDecodedModified() } } - if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) { - bodyModified = true - disablePatch() - } - - // Apply OpenAI fast policy (参照 Claude BetaPolicy 的 fast-mode 过滤): - // 针对 body 的 service_tier 字段("priority" 即 fast,"flex"),按策略 - // 执行 filter(删除字段)或 block(拒绝请求)。对 gpt-5.5 等模型屏蔽 - // fast 时在此生效。 - // - // 注意: - // 1. 此处统一使用 upstreamModel(已经过 GetMappedModel + - // normalizeOpenAIModelForUpstream + Codex OAuth normalize),与 - // chat-completions / messages 入口保持一致,避免不同入口因为模型 - // 维度不同而出现 whitelist 命中差异。 - // 2. action=pass 时也要把 raw "fast" 归一化为 "priority" 写回 body, - // 否则 native /responses 入口透传 "fast" 给上游会被拒。chat- - // completions 入口由 normalizeResponsesBodyServiceTier 完成同一 - // 行为,这里手工实现等效逻辑。 - if rawTier, ok := reqBody["service_tier"].(string); ok { + if rawTier := requestView.ServiceTier; rawTier != "" { if normTier := normalizedOpenAIServiceTierValue(rawTier); normTier != "" { action, errMsg := s.evaluateOpenAIFastPolicy(ctx, account, upstreamModel, normTier) switch action { @@ -2738,46 +2646,51 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco writeOpenAIFastPolicyBlockedResponse(c, blocked) return nil, blocked case BetaPolicyActionFilter: - delete(reqBody, "service_tier") - bodyModified = true - disablePatch() + markPatchDelete("service_tier") default: - // pass:若客户端传的是别名 "fast",归一化为 "priority" - // 后写回 body,确保上游收到的是其能识别的规范值。 if normTier != rawTier { - reqBody["service_tier"] = normTier - bodyModified = true markPatchSet("service_tier", normTier) } } } } - if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed { - MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) - c.JSON(http.StatusForbidden, gin.H{ - "error": gin.H{ - "type": "permission_error", - "message": ImageGenerationPermissionMessage(), - }, - }) - return nil, errors.New("image generation disabled for group") + if bodyModified { + if requestView.HasPatches() { + if patchedBody, patchErr := requestView.ApplyPatches(); patchErr == nil { + body = patchedBody + requestView = newOpenAIRequestView(body) + reqBody = nil + bodyModified = false + } + } + if bodyModified { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + var marshalErr error + body, marshalErr = marshalOpenAIUpstreamJSON(decoded) + if marshalErr != nil { + return nil, fmt.Errorf("serialize request body: %w", marshalErr) + } + requestView = newOpenAIRequestView(body) + } } imageBillingModel := "" imageSizeTier := "" imageInputSize := "" - if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) { + if imageIntent { + var imageCfg OpenAIResponsesImageBillingConfig var imageCfgErr error - imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel) + if reqBody != nil { + imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel) + } else { + imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, billingModel) + } if imageCfgErr != nil { setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "") - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": imageCfgErr.Error(), - "param": "size", - }, - }) + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": imageCfgErr.Error(), "param": "size"}}) return nil, imageCfgErr } imageBillingModel = imageCfg.Model @@ -2785,29 +2698,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco imageInputSize = imageCfg.InputSize } - // Re-serialize body only if modified - if bodyModified { - serializedByPatch := false - if !patchDisabled && patchHasOp { - var patchErr error - if patchDelete { - body, patchErr = sjson.DeleteBytes(body, patchPath) - } else { - body, patchErr = sjson.SetBytes(body, patchPath, patchValue) - } - if patchErr == nil { - serializedByPatch = true - } - } - if !serializedByPatch { - var marshalErr error - body, marshalErr = marshalOpenAIUpstreamJSON(reqBody) - if marshalErr != nil { - return nil, fmt.Errorf("serialize request body: %w", marshalErr) - } - } - } - // Get access token token, _, err := s.GetAccessToken(ctx, account) if err != nil { @@ -2816,8 +2706,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco // 命中 WS 时仅走 WebSocket Mode;不再自动回退 HTTP。 if wsDecision.Transport == OpenAIUpstreamTransportResponsesWebsocketV2 { - // WS 分支不会再回落 HTTP;重连恢复可直接更新 reqBody,避免额外保留一份完整顶层 map。 - wsReqBody := reqBody + // WS 分支需要结构化 payload 与重连恢复,命中后再触发 full-map decode。 + wsReqBody, err := ensureReqBody() + if err != nil { + return nil, err + } _, hasPreviousResponseID := wsReqBody["previous_response_id"] logOpenAIWSModeDebug( "forward_start account_id=%d account_type=%s model=%s stream=%v has_previous_response_id=%v", @@ -3076,8 +2969,12 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) upstreamCode := extractUpstreamErrorCode(respBody) if !httpInvalidEncryptedContentRetryTried && resp.StatusCode == http.StatusBadRequest && upstreamCode == "invalid_encrypted_content" { - if trimOpenAIEncryptedReasoningItems(reqBody) { - body, err = marshalOpenAIUpstreamJSON(reqBody) + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if trimOpenAIEncryptedReasoningItems(decoded) { + body, err = marshalOpenAIUpstreamJSON(decoded) if err != nil { return nil, fmt.Errorf("serialize invalid_encrypted_content retry body: %w", err) } @@ -3118,8 +3015,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } defer func() { _ = resp.Body.Close() }() - reasoningEffort := extractOpenAIReasoningEffort(reqBody, originalModel) - serviceTier := extractOpenAIServiceTier(reqBody) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel) + serviceTier := extractOpenAIServiceTierFromBody(body) // 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。 reqBody = nil @@ -6283,6 +6180,14 @@ type openAIRequestView struct { PreviousResponseID string ServiceTier string ReasoningEffort string + patches []openAIRequestPatch + patchesDisabled bool +} + +type openAIRequestPatch struct { + path string + delete bool + value any } func newOpenAIRequestView(body []byte) openAIRequestView { @@ -6305,6 +6210,110 @@ func (v openAIRequestView) Decode(c *gin.Context) (map[string]any, error) { return getOpenAIRequestBodyMap(c, v.body) } +func (v *openAIRequestView) MarkPatchSet(path string, value any) { + if v == nil || v.patchesDisabled { + return + } + path = strings.TrimSpace(path) + if path == "" { + v.DisablePatches() + return + } + v.patches = append(v.patches, openAIRequestPatch{path: path, value: value}) +} + +func (v *openAIRequestView) MarkPatchDelete(path string) { + if v == nil || v.patchesDisabled { + return + } + path = strings.TrimSpace(path) + if path == "" { + v.DisablePatches() + return + } + v.patches = append(v.patches, openAIRequestPatch{path: path, delete: true}) +} + +func (v *openAIRequestView) DisablePatches() { + if v == nil { + return + } + v.patchesDisabled = true + v.patches = nil +} + +func (v openAIRequestView) HasPatches() bool { + return !v.patchesDisabled && len(v.patches) > 0 +} + +func (v openAIRequestView) ApplyPatches() ([]byte, error) { + if v.patchesDisabled || len(v.patches) == 0 { + return nil, errors.New("openai request patches disabled") + } + body := v.body + for _, patch := range v.patches { + var err error + if patch.delete { + body, err = sjson.DeleteBytes(body, patch.path) + } else { + body, err = sjson.SetBytes(body, patch.path, patch.value) + } + if err != nil { + return nil, err + } + } + return body, nil +} + +func setOpenAIRequestMapPath(reqBody map[string]any, path string, value any) { + path = strings.TrimSpace(path) + if reqBody == nil || path == "" { + return + } + parts := strings.Split(path, ".") + current := reqBody + for _, part := range parts[:len(parts)-1] { + part = strings.TrimSpace(part) + if part == "" { + return + } + next, _ := current[part].(map[string]any) + if next == nil { + next = map[string]any{} + current[part] = next + } + current = next + } + last := strings.TrimSpace(parts[len(parts)-1]) + if last != "" { + current[last] = value + } +} + +func deleteOpenAIRequestMapPath(reqBody map[string]any, path string) { + path = strings.TrimSpace(path) + if reqBody == nil || path == "" { + return + } + parts := strings.Split(path, ".") + current := reqBody + for _, part := range parts[:len(parts)-1] { + part = strings.TrimSpace(part) + if part == "" { + return + } + next, _ := current[part].(map[string]any) + if next == nil { + return + } + current = next + } + last := strings.TrimSpace(parts[len(parts)-1]) + if last != "" { + delete(current, last) + } +} + func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) { view := newOpenAIRequestView(body) return view.Model, view.Stream, view.PromptCacheKey @@ -6758,8 +6767,84 @@ func buildOpenAIFastPolicyBlockedWSEvent(err *OpenAIFastBlockedError) []byte { return payload } +func openAIRequestBodyMayContainImageInput(body []byte) bool { + if len(body) == 0 { + return false + } + input := gjson.GetBytes(body, "input") + messages := gjson.GetBytes(body, "messages") + return openAIJSONValueMayContainImageInput(input) || openAIJSONValueMayContainImageInput(messages) +} + +func openAIJSONValueMayContainImageInput(value gjson.Result) bool { + if !value.Exists() { + return false + } + if value.IsArray() { + found := false + value.ForEach(func(_, item gjson.Result) bool { + if openAIJSONValueMayContainImageInput(item) { + found = true + return false + } + return true + }) + return found + } + if value.IsObject() { + if strings.TrimSpace(value.Get("type").String()) == "input_image" || value.Get("image_url").Exists() { + return true + } + return openAIJSONValueMayContainImageInput(value.Get("content")) + } + return false +} + +func openAIRequestBodyMayContainEmptyBase64InputImage(body []byte) bool { + if len(body) == 0 || !openAIRequestBodyMayContainInputImageToken(body) { + return false + } + input := gjson.GetBytes(body, "input") + if !input.Exists() { + return false + } + return openAIJSONValueMayContainEmptyBase64InputImage(input) +} + +func openAIRequestBodyMayContainInputImageToken(body []byte) bool { + if bytes.Contains(body, []byte("input_image")) { + return true + } + // JSON 字符串任意字符都可能被 unicode escape,遇到 \u 时交给 gjson 解码后的结构扫描兜底。 + return bytes.Contains(body, []byte("\\u")) +} + +func openAIJSONValueMayContainEmptyBase64InputImage(value gjson.Result) bool { + if !value.Exists() { + return false + } + if value.IsArray() { + found := false + value.ForEach(func(_, item gjson.Result) bool { + if openAIJSONValueMayContainEmptyBase64InputImage(item) { + found = true + return false + } + return true + }) + return found + } + if value.IsObject() { + if strings.TrimSpace(value.Get("type").String()) == "input_image" && isEmptyBase64DataURI(value.Get("image_url").String()) { + return true + } + return openAIJSONValueMayContainEmptyBase64InputImage(value.Get("content")) + } + return false +} + func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, error) { - if len(body) == 0 || !bytes.Contains(body, []byte(`"image_url"`)) || !bytes.Contains(body, []byte(`base64,`)) { + if !openAIRequestBodyMayContainEmptyBase64InputImage(body) { return body, false, nil } diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index df17ed2e..d240cec8 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -1,12 +1,18 @@ package service import ( + "context" "encoding/json" + "io" + "net/http" "net/http/httptest" + "strings" "testing" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" ) func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) { @@ -29,6 +35,454 @@ func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) { require.IsType(t, []any{}, reqBody["input"]) } +func TestOpenAIRequestView_ApplyPatches(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5","previous_response_id":"resp_1","reasoning":{"effort":"minimal"},"input":[{"type":"message","content":"hi"}]}`)) + view.MarkPatchSet("model", "gpt-5.1") + view.MarkPatchDelete("previous_response_id") + view.MarkPatchSet("reasoning.effort", "none") + + patched, err := view.ApplyPatches() + require.NoError(t, err) + require.JSONEq(t, `{"model":"gpt-5.1","reasoning":{"effort":"none"},"input":[{"type":"message","content":"hi"}]}`, string(patched)) +} + +func TestOpenAIRequestView_ApplyPatchesDisabled(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`)) + view.MarkPatchSet("model", "gpt-5.1") + view.DisablePatches() + + _, err := view.ApplyPatches() + require.Error(t, err) +} + +func TestOpenAIRequestView_HasPatches(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`)) + require.False(t, view.HasPatches()) + + view.MarkPatchSet("model", "gpt-5.1") + require.True(t, view.HasPatches()) + + view.DisablePatches() + require.False(t, view.HasPatches()) +} + +func TestOpenAIGatewayService_Forward_HTTPPatchPathKeepsLargeInputRaw(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"usage":{"input_tokens":1,"output_tokens":2,"input_tokens_details":{"cached_tokens":0}}}`, + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 1, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5","stream":false,"reasoning":{"effort":"minimal"},"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.JSONEq(t, `{"model":"gpt-5","stream":false,"reasoning":{"effort":"none"},"instructions":"You are a helpful coding assistant.","input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`, string(upstream.lastBody)) + require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw) +} + +func TestOpenAIGatewayService_Forward_DecodedMutationKeepsLaterFieldDeletes(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 2, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation","format":"png"}],"input":[{"type":"message","content":"draw"}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.False(t, gjson.GetBytes(upstream.lastBody, "max_completion_tokens").Exists()) + require.False(t, gjson.GetBytes(upstream.lastBody, "tools.0.format").Exists()) + require.Equal(t, "png", gjson.GetBytes(upstream.lastBody, "tools.0.output_format").String()) +} + +func TestOpenAIGatewayService_Forward_MappedImageModelUsesImageGate(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 3, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + "model_mapping": map[string]any{"draw-alias": "gpt-image-2"}, + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: false}}) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"draw-alias","stream":false,"input":"draw"}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.Error(t, err) + require.Nil(t, result) + require.Nil(t, upstream.lastReq) + require.Equal(t, http.StatusForbidden, rec.Code) +} + +func TestOpenAIGatewayService_Forward_TextDataImageDoesNotForceMapMarshal(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 4, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"message","content":[{"type":"input_text","text":"literal data:image/png;base64, only","nonce":1e1000000}]}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw) +} + +func TestOpenAIGatewayService_Forward_ImageToolBillingDoesNotForceFullDecode(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image"}],"usage":{"input_tokens":1,"output_tokens":2}}`, + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 9, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"draw","nonce":1e1000000}]}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw) + require.Equal(t, 1, result.ImageCount) + require.Equal(t, "2K", result.ImageSize) + require.Equal(t, "gpt-image-2", result.BillingModel) +} + +func TestOpenAIGatewayService_Forward_HTTPRetryRecoveryDoesNotDecodeBeforeError(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + responses: []*http.Response{ + { + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_encrypted_content","type":"invalid_request_error","message":"bad encrypted content"}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 10, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"reasoning","encrypted_content":"gAAA","summary":[{"type":"summary_text","text":"keep me"}]},{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 2) + require.Equal(t, "gAAA", gjson.GetBytes(upstream.bodies[0], "input.0.encrypted_content").String()) + require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.bodies[0], "input.1.content.0.nonce").Raw) + require.False(t, gjson.GetBytes(upstream.bodies[1], "input.0.encrypted_content").Exists()) + require.Equal(t, "summary_text", gjson.GetBytes(upstream.bodies[1], "input.0.summary.0.type").String()) +} + +func TestOpenAIGatewayService_Forward_CodexSparkRejectsEscapedInputImage(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 5, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.3-codex-spark","stream":false,"input":[{"type":"input_` + "\\u0069" + `mage","file_id":"file_1"}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.Error(t, err) + require.Nil(t, result) + require.Nil(t, upstream.lastReq) + require.Equal(t, http.StatusBadRequest, rec.Code) +} + +func TestOpenAIGatewayService_Forward_CodexBridgeInjectionSetsImageBilling(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image","size":"1024x1024"}],"usage":{"input_tokens":1,"output_tokens":2}}`, + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Gateway.ForceCodexCLI = true + cfg.Gateway.CodexImageGenerationBridgeEnabled = true + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 7, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: true}}) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5","stream":false,"input":"draw if needed"}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, result.ImageCount) + require.Equal(t, "2K", result.ImageSize) + require.Equal(t, "gpt-image-2", result.BillingModel) +} + +func TestOpenAIGatewayService_Forward_HTTPDeletesPreviousResponseIDWhenPresent(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + account := &Account{ + ID: 8, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + + for _, body := range [][]byte{ + []byte(`{"model":"gpt-5","stream":false,"previous_response_id":"","input":"hi"}`), + []byte(`{"model":"gpt-5","stream":false,"previous_response_id":null,"input":"hi"}`), + } { + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.False(t, gjson.GetBytes(upstream.lastBody, "previous_response_id").Exists()) + } +} + +func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedJSON(t *testing.T) { + body := []byte(`{"input":[{"type":"message","content":[{"type":"input_image","image_` + "\\u0075" + `rl":"data:image/png;base64` + "\\u002c" + ` "}]}]}`) + + require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body)) +} + +func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedImageType(t *testing.T) { + body := []byte(`{"input":[{"type":"message","content":[{"type":"input_` + "\\u0069" + `mage","image_url":"data:image/png;base64, "}]}]}`) + + require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body)) +} + +func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedInputPrefix(t *testing.T) { + body := []byte(`{"input":[{"type":"message","content":[{"type":"inp` + "\\u0075" + `t_image","image_url":"data:image/png;base64, "}]}]}`) + + require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body)) +} + +func TestOpenAIGatewayService_Forward_ImageOnlyModelKeepsSupportedVerbosity(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 6, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-image-2","stream":false,"text":{"verbosity":"low"},"input":"draw"}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "low", gjson.GetBytes(upstream.lastBody, "text.verbosity").String()) + require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String()) +} + func TestExtractOpenAIRequestMetaFromBody(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index e949560f..bd262207 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -171,6 +171,91 @@ func TestOpenAIGatewayService_Forward_WSv2_SuccessAndBindSticky(t *testing.T) { require.Equal(t, "resp_new_1", gjson.GetBytes(responseBody, "id").String()) } +func TestOpenAIGatewayService_Forward_WSv2_UsesPatchedBodyAfterValidationDecode(t *testing.T) { + gin.SetMode(gin.TestMode) + + type receivedPayload struct { + MaxCompletionTokensExists bool + } + receivedCh := make(chan receivedPayload, 1) + + upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }} + wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade websocket failed: %v", err) + return + } + defer func() { _ = conn.Close() }() + + var request map[string]any + if err := conn.ReadJSON(&request); err != nil { + t.Errorf("read ws request failed: %v", err) + return + } + requestJSON := requestToJSONString(request) + receivedCh <- receivedPayload{MaxCompletionTokensExists: gjson.Get(requestJSON, "max_completion_tokens").Exists()} + + if err := conn.WriteJSON(map[string]any{ + "type": "response.completed", + "response": map[string]any{ + "id": "resp_patched_ws_1", + "model": "gpt-5.3-codex-spark", + "usage": map[string]any{"input_tokens": 1, "output_tokens": 1}, + }, + }); err != nil { + t.Errorf("write response.completed failed: %v", err) + return + } + })) + defer wsServer.Close() + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("User-Agent", "unit-test-agent/1.0") + + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 30 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 10 + + svc := &OpenAIGatewayService{ + cfg: cfg, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + } + + account := &Account{ + ID: 10, + Name: "openai-ws", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": wsServer.URL, + }, + Extra: map[string]any{"responses_websockets_v2_enabled": true}, + } + + body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation"}],"input":[{"type":"input_text","text":"hello"}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.OpenAIWSMode) + + received := <-receivedCh + require.False(t, received.MaxCompletionTokensExists) +} + func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testing.T) { gin.SetMode(gin.TestMode)