diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index fb05b64b01..8942404eaa 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" + "github.com/tidwall/sjson" ) type GrokMediaEndpoint string @@ -296,6 +297,10 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( if err != nil { return nil, err } + body, contentType, err = normalizeGrokMediaForwardBody(endpoint, body, contentType) + if err != nil { + return nil, err + } var bodyReader io.Reader if endpoint.RequiresRequestBody() { @@ -424,6 +429,33 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten return out, "application/json", nil } +func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, contentType string) ([]byte, string, error) { + if !endpoint.RequiresRequestBody() || !gjson.ValidBytes(body) { + return body, contentType, nil + } + model := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + upstreamModel := normalizeGrokMediaModelForEndpoint(endpoint, model) + if upstreamModel == "" || upstreamModel == model { + return body, contentType, nil + } + out, err := sjson.SetBytes(body, "model", upstreamModel) + if err != nil { + return nil, "", fmt.Errorf("rewrite grok media model: %w", err) + } + return out, contentType, nil +} + +func normalizeGrokMediaModelForEndpoint(endpoint GrokMediaEndpoint, model string) string { + model = strings.TrimSpace(model) + switch endpoint { + case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits: + if model == "grok-imagine" { + return "grok-imagine-image-quality" + } + } + return model +} + type grokMediaUsageMetadata struct { ResponseID string Usage OpenAIUsage diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index feb2c95abf..eae424ce9a 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -201,7 +201,28 @@ func TestParseGrokMediaRequestBuildsMultipartModerationBody(t *testing.T) { require.True(t, strings.HasPrefix(gjson.GetBytes(moderationBody, "images.0.image_url").String(), "data:image/")) } -func TestForwardGrokMediaImagesGenerationPassthrough(t *testing.T) { +func TestNormalizeGrokMediaModelForEndpoint(t *testing.T) { + tests := []struct { + name string + endpoint GrokMediaEndpoint + model string + want string + }{ + {name: "image generation alias", endpoint: GrokMediaEndpointImagesGenerations, model: "grok-imagine", want: "grok-imagine-image-quality"}, + {name: "image edit alias", endpoint: GrokMediaEndpointImagesEdits, model: "grok-imagine", want: "grok-imagine-image-quality"}, + {name: "image quality passthrough", endpoint: GrokMediaEndpointImagesGenerations, model: "grok-imagine-image-quality", want: "grok-imagine-image-quality"}, + {name: "image fast passthrough", endpoint: GrokMediaEndpointImagesGenerations, model: "grok-imagine-image", want: "grok-imagine-image"}, + {name: "video passthrough", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video", want: "grok-imagine-video"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, normalizeGrokMediaModelForEndpoint(tt.endpoint, tt.model)) + }) + } +} + +func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) @@ -238,12 +259,12 @@ func TestForwardGrokMediaImagesGenerationPassthrough(t *testing.T) { require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type")) - require.JSONEq(t, string(body), string(upstream.lastBody)) + require.JSONEq(t, `{"model":"grok-imagine-image-quality","prompt":"draw a cat"}`, string(upstream.lastBody)) require.Equal(t, http.StatusOK, recorder.Code) require.JSONEq(t, `{"data":[]}`, recorder.Body.String()) require.Equal(t, "xai-image-req", result.RequestID) - require.Equal(t, "grok-imagine", result.Model) - require.Equal(t, "grok-imagine", result.BillingModel) + require.Equal(t, "grok-imagine-image-quality", result.Model) + require.Equal(t, "grok-imagine-image-quality", result.BillingModel) require.Equal(t, 1, result.ImageCount) require.Equal(t, ImageBillingSize2K, result.ImageSize) }