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