fix: harden grok media routing

This commit is contained in:
Heatherm Huang
2026-07-01 15:36:08 +08:00
parent a34d4967e6
commit 42e471f59a
4 changed files with 492 additions and 60 deletions
+15 -10
View File
@@ -14,7 +14,6 @@ import (
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"go.uber.org/zap"
)
@@ -84,7 +83,8 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
}
contentType := c.GetHeader("Content-Type")
requestModel := service.ExtractGrokMediaModel(contentType, body)
requestInfo := service.ParseGrokMediaRequest(contentType, body)
requestModel := requestInfo.Model
if endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) == "" {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
@@ -103,7 +103,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
return
}
if moderationBody := grokMediaModerationBody(body); len(moderationBody) > 0 {
if moderationBody := requestInfo.ModerationBody(); len(moderationBody) > 0 {
decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, moderationBody)
if decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
@@ -149,6 +149,9 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
sessionSeed = []byte(requestID)
}
sessionHash := h.gatewayService.GenerateExplicitSessionHash(c, sessionSeed)
if endpoint == service.GrokMediaEndpointVideoStatus {
sessionHash = service.GrokMediaVideoRequestSessionHash(requestID)
}
requestCtx := c.Request.Context()
failedAccountIDs := make(map[int64]struct{})
sameAccountRetryCount := make(map[int64]int)
@@ -294,6 +297,15 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
if endpoint == service.GrokMediaEndpointVideosGenerations && strings.TrimSpace(result.ResponseID) != "" {
if err := h.gatewayService.BindGrokMediaVideoRequestAccount(requestCtx, apiKey.GroupID, result.ResponseID, account.ID); err != nil {
reqLog.Warn("grok_media.bind_video_request_account_failed",
zap.Int64("account_id", account.ID),
zap.String("request_id", result.ResponseID),
zap.Error(err),
)
}
}
if shouldRecordGrokMediaUsage(endpoint, requestModel) {
recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, result, requestModel, body, requestID)
}
@@ -305,13 +317,6 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
}
}
func grokMediaModerationBody(body []byte) []byte {
if gjson.ValidBytes(body) {
return body
}
return nil
}
func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string) bool {
return endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) != ""
}
+1 -1
View File
@@ -177,7 +177,7 @@ func forwardLLMTesterRequest(c *gin.Context, req llmTesterProxyRequest, method,
response.Error(c, http.StatusBadGateway, fmt.Sprintf("upstream request failed: %s", err.Error()))
return
}
defer upstreamResp.Body.Close()
defer func() { _ = upstreamResp.Body.Close() }()
payload, err := readLLMTesterResponseBody(upstreamResp.Body)
if err != nil {
+351 -49
View File
@@ -3,11 +3,13 @@ package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime"
"mime/multipart"
"net/http"
"strconv"
"strings"
"time"
@@ -39,6 +41,56 @@ func (e GrokMediaEndpoint) IsGenerationRequest() bool {
}
}
type GrokMediaRequestInfo struct {
Model string
Prompt string
N int
Size string
SizeTier string
InputImageURLs []string
MaskImageURL string
Uploads []OpenAIImagesUpload
MaskUpload *OpenAIImagesUpload
}
func (r GrokMediaRequestInfo) ModerationBody() []byte {
payload := map[string]any{}
if prompt := strings.TrimSpace(r.Prompt); prompt != "" {
payload["prompt"] = prompt
}
images := make([]map[string]string, 0, len(r.InputImageURLs)+len(r.Uploads)+1)
for _, imageURL := range r.InputImageURLs {
if imageURL = strings.TrimSpace(imageURL); imageURL != "" {
images = append(images, map[string]string{"image_url": imageURL})
}
}
for _, upload := range r.Uploads {
if dataURL := upload.ModerationDataURL(); dataURL != "" {
images = append(images, map[string]string{"image_url": dataURL})
}
}
if maskURL := strings.TrimSpace(r.MaskImageURL); maskURL != "" {
images = append(images, map[string]string{"image_url": maskURL})
}
if r.MaskUpload != nil {
if dataURL := r.MaskUpload.ModerationDataURL(); dataURL != "" {
images = append(images, map[string]string{"image_url": dataURL})
}
}
if len(images) > 0 {
payload["images"] = images
}
if len(payload) == 0 {
return nil
}
body, err := json.Marshal(payload)
if err != nil {
return nil
}
return body
}
func (e GrokMediaEndpoint) httpMethod() string {
if e == GrokMediaEndpointVideoStatus {
return http.MethodGet
@@ -47,41 +99,158 @@ func (e GrokMediaEndpoint) httpMethod() string {
}
func ExtractGrokMediaModel(contentType string, body []byte) string {
if model := strings.TrimSpace(gjson.GetBytes(body, "model").String()); model != "" {
return model
}
return extractGrokMediaMultipartModel(contentType, body)
return ParseGrokMediaRequest(contentType, body).Model
}
func extractGrokMediaMultipartModel(contentType string, body []byte) string {
func ParseGrokMediaRequest(contentType string, body []byte) GrokMediaRequestInfo {
info := GrokMediaRequestInfo{N: 1}
if gjson.ValidBytes(body) {
parseGrokMediaJSONRequest(body, &info)
} else {
parseGrokMediaMultipartRequest(contentType, body, &info)
}
info.Model = strings.TrimSpace(info.Model)
info.Prompt = strings.TrimSpace(info.Prompt)
info.Size = strings.TrimSpace(info.Size)
info.SizeTier = NormalizeImageBillingTierOrDefault(info.Size)
if info.N <= 0 {
info.N = 1
}
return info
}
func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) {
if info == nil {
return
}
info.Model = strings.TrimSpace(gjson.GetBytes(body, "model").String())
info.Prompt = strings.TrimSpace(gjson.GetBytes(body, "prompt").String())
info.Size = strings.TrimSpace(gjson.GetBytes(body, "size").String())
if n := gjson.GetBytes(body, "n"); n.Exists() && n.Type == gjson.Number {
info.N = int(n.Int())
}
appendJSONImageURLs := func(value gjson.Result) {
if !value.Exists() {
return
}
switch {
case value.IsArray():
for _, item := range value.Array() {
if imageURL := strings.TrimSpace(item.Get("image_url").String()); imageURL != "" {
info.InputImageURLs = append(info.InputImageURLs, imageURL)
continue
}
if item.Type == gjson.String {
imageURL := strings.TrimSpace(item.String())
if imageURL == "" {
continue
}
info.InputImageURLs = append(info.InputImageURLs, imageURL)
}
}
default:
if imageURL := strings.TrimSpace(value.Get("image_url").String()); imageURL != "" {
info.InputImageURLs = append(info.InputImageURLs, imageURL)
return
}
if value.Type == gjson.String {
imageURL := strings.TrimSpace(value.String())
if imageURL == "" {
return
}
info.InputImageURLs = append(info.InputImageURLs, imageURL)
}
}
}
appendJSONImageURLs(gjson.GetBytes(body, "image"))
appendJSONImageURLs(gjson.GetBytes(body, "images"))
info.MaskImageURL = strings.TrimSpace(gjson.GetBytes(body, "mask.image_url").String())
}
func parseGrokMediaMultipartRequest(contentType string, body []byte, info *GrokMediaRequestInfo) {
if info == nil {
return
}
mediaType, params, err := mime.ParseMediaType(strings.TrimSpace(contentType))
if err != nil || !strings.EqualFold(mediaType, "multipart/form-data") {
return ""
return
}
boundary := strings.TrimSpace(params["boundary"])
if boundary == "" {
return ""
return
}
reader := multipart.NewReader(bytes.NewReader(body), boundary)
for {
part, err := reader.NextPart()
if err == io.EOF {
return ""
return
}
if err != nil {
return ""
return
}
if part.FormName() != "model" || part.FileName() != "" {
name := strings.TrimSpace(part.FormName())
if name == "" {
_ = part.Close()
continue
}
data, err := io.ReadAll(part)
data, err := io.ReadAll(io.LimitReader(part, openAIImageMaxUploadPartSize))
_ = part.Close()
if err != nil {
return ""
return
}
fileName := strings.TrimSpace(part.FileName())
partContentType := strings.TrimSpace(part.Header.Get("Content-Type"))
if fileName != "" {
upload := OpenAIImagesUpload{
FieldName: name,
FileName: fileName,
ContentType: partContentType,
Data: data,
}
if name == "mask" {
info.MaskUpload = &upload
continue
}
if name == "image" || strings.HasPrefix(name, "image[") {
info.Uploads = append(info.Uploads, upload)
}
continue
}
value := strings.TrimSpace(string(data))
switch name {
case "model":
info.Model = value
case "prompt":
info.Prompt = value
case "size":
info.Size = value
case "n":
if n, err := strconv.Atoi(value); err == nil {
info.N = n
}
case "image", "image_url":
if value != "" {
info.InputImageURLs = append(info.InputImageURLs, value)
}
case "mask", "mask_image_url":
info.MaskImageURL = value
}
return strings.TrimSpace(string(data))
}
}
func GrokMediaVideoRequestSessionHash(requestID string) string {
requestID = strings.TrimSpace(requestID)
if requestID == "" {
return ""
}
return "grok-video:" + DeriveSessionHashFromSeed(requestID)
}
func (s *OpenAIGatewayService) BindGrokMediaVideoRequestAccount(ctx context.Context, groupID *int64, requestID string, accountID int64) error {
return s.BindStickySession(ctx, groupID, GrokMediaVideoRequestSessionHash(requestID), accountID)
}
func (e GrokMediaEndpoint) upstreamURL(baseURL, requestID string) (string, error) {
switch e {
case GrokMediaEndpointImagesGenerations:
@@ -157,39 +326,11 @@ func (s *OpenAIGatewayService) ForwardGrokMedia(
defer func() { _ = resp.Body.Close() }()
requestIDHeader := firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id"))
requestModel := ExtractGrokMediaModel(contentType, body)
requestInfo := ParseGrokMediaRequest(contentType, body)
requestModel := requestInfo.Model
if resp.StatusCode >= 400 {
respBody := s.readUpstreamErrorBody(resp)
s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode))
upstreamMsg := sanitizeUpstreamErrorMessage(extractUpstreamErrorMessage(respBody))
if upstreamMsg == "" {
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: requestIDHeader,
Kind: "failover",
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverUpstreamError(resp.StatusCode) {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
writeGrokMediaResponse(c, resp, respBody, s.responseHeaderFilter)
return &OpenAIForwardResult{
RequestID: requestIDHeader,
Model: requestModel,
UpstreamModel: requestModel,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
}, nil
return s.handleGrokMediaErrorResponse(ctx, resp, c, account, requestIDHeader, requestModel)
}
s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode))
@@ -198,15 +339,176 @@ func (s *OpenAIGatewayService) ForwardGrokMedia(
return nil, err
}
writeGrokMediaResponse(c, resp, respBody, s.responseHeaderFilter)
usage := grokMediaUsageFromResponse(endpoint, requestInfo, respBody)
return &OpenAIForwardResult{
RequestID: requestIDHeader,
Model: requestModel,
UpstreamModel: requestModel,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
RequestID: requestIDHeader,
ResponseID: usage.ResponseID,
Usage: usage.Usage,
Model: requestModel,
BillingModel: requestModel,
UpstreamModel: requestModel,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
ImageCount: usage.ImageCount,
ImageSize: usage.ImageSize,
ImageInputSize: usage.ImageInputSize,
ImageOutputSizes: usage.ImageOutputSizes,
}, nil
}
type grokMediaUsageMetadata struct {
ResponseID string
Usage OpenAIUsage
ImageCount int
ImageSize string
ImageInputSize string
ImageOutputSizes []string
}
func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMediaRequestInfo, responseBody []byte) grokMediaUsageMetadata {
usage, _ := extractOpenAIUsageFromJSONBytes(responseBody)
meta := grokMediaUsageMetadata{Usage: usage}
switch endpoint {
case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits:
imageCount := countOpenAIResponseImageOutputsFromJSONBytes(responseBody)
if imageCount <= 0 {
imageCount = requestInfo.N
}
if imageCount <= 0 {
imageCount = 1
}
meta.ImageCount = imageCount
meta.ImageSize = requestInfo.SizeTier
meta.ImageInputSize = requestInfo.Size
meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody)
case GrokMediaEndpointVideosGenerations:
meta.ResponseID = extractGrokMediaVideoRequestID(responseBody)
meta.ImageCount = 1
meta.ImageSize = requestInfo.SizeTier
meta.ImageInputSize = requestInfo.Size
}
return meta
}
func extractGrokMediaVideoRequestID(body []byte) string {
if len(body) == 0 || !gjson.ValidBytes(body) {
return ""
}
for _, path := range []string{"request_id", "id", "data.request_id", "data.id", "video.request_id", "video.id"} {
if id := strings.TrimSpace(gjson.GetBytes(body, path).String()); id != "" {
return id
}
}
return ""
}
func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
ctx context.Context,
resp *http.Response,
c *gin.Context,
account *Account,
requestIDHeader string,
requestedModel string,
) (*OpenAIForwardResult, error) {
body := s.readUpstreamErrorBody(resp)
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(body)))
if upstreamMsg == "" {
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
}
upstreamDetail := ""
if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody {
maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes
if maxBytes <= 0 {
maxBytes = 2048
}
upstreamDetail = truncateString(string(body), maxBytes)
}
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
if status, errType, errMsg, matched := applyErrorPassthroughRule(
c,
account.Platform,
resp.StatusCode,
body,
http.StatusBadGateway,
"upstream_error",
"Upstream request failed",
); matched {
MarkResponseCommitted(c)
writeGrokMediaErrorResponse(c, status, errType, errMsg)
return nil, fmt.Errorf("upstream error: %d (passthrough rule matched) message=%s", resp.StatusCode, upstreamMsg)
}
if !account.ShouldHandleErrorCode(resp.StatusCode) {
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: requestIDHeader,
Kind: "http_error",
Message: upstreamMsg,
Detail: upstreamDetail,
})
MarkResponseCommitted(c)
writeGrokMediaErrorResponse(c, http.StatusInternalServerError, "upstream_error", "Upstream gateway error")
return nil, fmt.Errorf("upstream error: %d (not in custom error codes) message=%s", resp.StatusCode, upstreamMsg)
}
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
kind := "http_error"
if s.shouldFailoverUpstreamError(resp.StatusCode) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: requestIDHeader,
Kind: kind,
Message: upstreamMsg,
Detail: upstreamDetail,
})
if kind == "failover" {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: body,
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
MarkResponseCommitted(c)
writeGrokMediaErrorResponse(c, resp.StatusCode, grokMediaErrorType(resp.StatusCode), upstreamMsg)
return nil, fmt.Errorf("upstream error: %d %s", resp.StatusCode, upstreamMsg)
}
func grokMediaErrorType(statusCode int) string {
switch {
case statusCode == http.StatusBadRequest:
return "invalid_request_error"
case statusCode == http.StatusNotFound:
return "not_found_error"
case statusCode == http.StatusTooManyRequests:
return "rate_limit_error"
default:
return "upstream_error"
}
}
func writeGrokMediaErrorResponse(c *gin.Context, statusCode int, errType, message string) {
if c == nil || c.Writer == nil || c.Writer.Written() {
return
}
c.JSON(statusCode, gin.H{
"error": gin.H{
"type": strings.TrimSpace(errType),
"message": strings.TrimSpace(message),
},
})
}
func writeGrokMediaResponse(c *gin.Context, resp *http.Response, body []byte, filter *responseheaders.CompiledHeaderFilter) {
if c == nil || resp == nil {
return
@@ -10,6 +10,7 @@ import (
"mime/multipart"
"net/http"
"net/http/httptest"
"net/textproto"
"strings"
"testing"
"time"
@@ -157,6 +158,30 @@ func TestExtractGrokMediaModelSupportsJSONAndMultipart(t *testing.T) {
require.Equal(t, "grok-imagine-edit", ExtractGrokMediaModel(writer.FormDataContentType(), buf.Bytes()))
}
func TestParseGrokMediaRequestBuildsMultipartModerationBody(t *testing.T) {
var buf bytes.Buffer
writer := multipart.NewWriter(&buf)
require.NoError(t, writer.WriteField("prompt", "edit this private image"))
require.NoError(t, writer.WriteField("model", "grok-imagine-edit"))
partHeader := textproto.MIMEHeader{}
partHeader.Set("Content-Disposition", `form-data; name="image"; filename="input.png"`)
partHeader.Set("Content-Type", "image/png")
part, err := writer.CreatePart(partHeader)
require.NoError(t, err)
_, err = part.Write([]byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a})
require.NoError(t, err)
require.NoError(t, writer.Close())
info := ParseGrokMediaRequest(writer.FormDataContentType(), buf.Bytes())
require.Equal(t, "grok-imagine-edit", info.Model)
require.Equal(t, "edit this private image", info.Prompt)
moderationBody := info.ModerationBody()
require.NotEmpty(t, moderationBody)
require.Equal(t, "edit this private image", gjson.GetBytes(moderationBody, "prompt").String())
require.True(t, strings.HasPrefix(gjson.GetBytes(moderationBody, "images.0.image_url").String(), "data:image/"))
}
func TestForwardGrokMediaImagesGenerationPassthrough(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
@@ -199,6 +224,50 @@ func TestForwardGrokMediaImagesGenerationPassthrough(t *testing.T) {
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, 1, result.ImageCount)
require.Equal(t, ImageBillingSize2K, result.ImageSize)
}
func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(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":"waves"}`)
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"},
"Xai-Request-Id": []string{"xai-video-generate-req"},
},
Body: io.NopCloser(strings.NewReader(`{"request_id":"video-request-123","usage":{"prompt_tokens":3,"completion_tokens":4}}`)),
}}
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.Equal(t, "video-request-123", result.ResponseID)
require.Equal(t, "grok-imagine-video-1.5", result.BillingModel)
require.Equal(t, 3, result.Usage.InputTokens)
require.Equal(t, 4, result.Usage.OutputTokens)
require.Equal(t, 1, result.ImageCount)
}
func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) {
@@ -242,6 +311,62 @@ func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) {
require.Equal(t, "xai-video-req", result.RequestID)
}
func TestBindGrokMediaVideoRequestAccountUsesRequestIDStickyHash(t *testing.T) {
ctx := context.Background()
groupID := int64(7)
cache := &stubGatewayCache{}
svc := &OpenAIGatewayService{cache: cache}
hash := GrokMediaVideoRequestSessionHash("video-request-123")
require.NotEmpty(t, hash)
require.NoError(t, svc.BindGrokMediaVideoRequestAccount(ctx, &groupID, "video-request-123", 63))
accountID, err := svc.getStickySessionAccountID(ctx, &groupID, hash)
require.NoError(t, err)
require.Equal(t, int64(63), accountID)
}
func TestForwardGrokMediaErrorHonorsCustomErrorCodes(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok-imagine","prompt":"draw a cat"}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
account := &Account{
ID: 64,
Name: "grok",
Platform: PlatformGrok,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "api-key",
"base_url": "https://xai.test/v1",
"custom_error_codes_enabled": true,
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{
"Content-Type": []string{"application/json"},
"Xai-Request-Id": []string{"xai-error-req"},
},
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"do not expose this upstream detail"}}`)),
}}
svc := &OpenAIGatewayService{httpUpstream: upstream}
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesGenerations, "", body, "application/json")
require.Error(t, err)
require.Nil(t, result)
require.Equal(t, http.StatusInternalServerError, recorder.Code)
require.Contains(t, recorder.Body.String(), "Upstream gateway error")
require.NotContains(t, recorder.Body.String(), "do not expose")
}
func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *testing.T) {
gin.SetMode(gin.TestMode)