diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index f8b7ff3e83..09d70eb0bd 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -433,9 +433,9 @@ func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, cont 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 { + info := ParseGrokMediaRequest(contentType, body) + upstreamModel := normalizeGrokMediaModelForEndpoint(endpoint, info.Model, info.HasInputImage()) + if upstreamModel == "" || upstreamModel == info.Model { return body, contentType, nil } out, err := sjson.SetBytes(body, "model", upstreamModel) @@ -445,13 +445,21 @@ func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, cont return out, contentType, nil } -func normalizeGrokMediaModelForEndpoint(endpoint GrokMediaEndpoint, model string) string { +func (r GrokMediaRequestInfo) HasInputImage() bool { + return len(r.InputImageURLs) > 0 || len(r.Uploads) > 0 +} + +func normalizeGrokMediaModelForEndpoint(endpoint GrokMediaEndpoint, model string, hasInputImage bool) string { model = strings.TrimSpace(model) switch endpoint { case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits: if model == "grok-imagine" { return "grok-imagine-image-quality" } + case GrokMediaEndpointVideosGenerations: + if model == "grok-imagine-video-1.5" && !hasInputImage { + return "grok-imagine-video" + } } return model } diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index d012ab533a..f6aa4b6cd1 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -203,21 +203,24 @@ func TestParseGrokMediaRequestBuildsMultipartModerationBody(t *testing.T) { func TestNormalizeGrokMediaModelForEndpoint(t *testing.T) { tests := []struct { - name string - endpoint GrokMediaEndpoint - model string - want string + name string + endpoint GrokMediaEndpoint + model string + hasInputImage bool + 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"}, + {name: "video 1.5 text-only fallback", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", want: "grok-imagine-video"}, + {name: "video 1.5 image-to-video passthrough", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", hasInputImage: true, want: "grok-imagine-video-1.5"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - require.Equal(t, tt.want, normalizeGrokMediaModelForEndpoint(tt.endpoint, tt.model)) + require.Equal(t, tt.want, normalizeGrokMediaModelForEndpoint(tt.endpoint, tt.model, tt.hasInputImage)) }) } } @@ -355,13 +358,52 @@ func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(t *testing.T) result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") require.NoError(t, err) require.Equal(t, "https://xai.test/v1/videos/generations", upstream.lastReq.URL.String()) + require.JSONEq(t, `{"model":"grok-imagine-video","prompt":"waves"}`, string(upstream.lastBody)) require.Equal(t, "video-request-123", result.ResponseID) - require.Equal(t, "grok-imagine-video-1.5", result.BillingModel) + require.Equal(t, "grok-imagine-video", result.BillingModel) require.Equal(t, 3, result.Usage.InputTokens) require.Equal(t, 4, result.Usage.OutputTokens) require.Equal(t, 1, result.ImageCount) } +func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-imagine-video-1.5","prompt":"animate","image":{"image_url":"data:image/png;base64,aW1n"}}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/videos/generations", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 63, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "api-key", + "base_url": "https://xai.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + }, + Body: io.NopCloser(strings.NewReader(`{"request_id":"video-request-456"}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") + require.NoError(t, err) + require.Equal(t, "https://xai.test/v1/videos/generations", upstream.lastReq.URL.String()) + require.JSONEq(t, `{"model":"grok-imagine-video-1.5","prompt":"animate","image":{"image_url":"data:image/png;base64,aW1n"}}`, string(upstream.lastBody)) + require.Equal(t, "video-request-456", result.ResponseID) + require.Equal(t, "grok-imagine-video-1.5", result.BillingModel) +} + func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode)