From 42e471f59ad0ea5b5dfe1120eeea881f6a480f1b Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Wed, 1 Jul 2026 15:36:08 +0800 Subject: [PATCH] fix: harden grok media routing --- backend/internal/handler/grok_media.go | 25 +- backend/internal/server/routes/llm_tester.go | 2 +- backend/internal/service/grok_media.go | 400 +++++++++++++++--- .../service/openai_gateway_grok_test.go | 125 ++++++ 4 files changed, 492 insertions(+), 60 deletions(-) diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 96023f452e..8e236ea49f 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -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) != "" } diff --git a/backend/internal/server/routes/llm_tester.go b/backend/internal/server/routes/llm_tester.go index 684b86d729..dd347177ee 100644 --- a/backend/internal/server/routes/llm_tester.go +++ b/backend/internal/server/routes/llm_tester.go @@ -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 { diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 9269cacf11..3b76b9c274 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -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 diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 268d07f9ec..a095243f57 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -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)