fix: normalize grok imagine image alias

This commit is contained in:
Heatherm Huang
2026-07-02 15:04:40 +08:00
parent 9934bd2572
commit 99a8d8ad28
2 changed files with 57 additions and 4 deletions
+32
View File
@@ -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
@@ -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)
}