fix: route text-only grok video requests

This commit is contained in:
Heatherm Huang
2026-07-07 16:44:20 +08:00
parent 44ab690a01
commit 3b2099350e
2 changed files with 60 additions and 10 deletions
+12 -4
View File
@@ -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
}
@@ -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)