mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix: route text-only grok video requests
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user