fix: route grok media endpoints

This commit is contained in:
Heatherm Huang
2026-07-01 11:43:35 +08:00
parent 5f62625ae6
commit 3b5d812f7a
9 changed files with 858 additions and 53 deletions
+7 -1
View File
@@ -21,6 +21,8 @@ const (
EndpointResponses = "/v1/responses"
EndpointImagesGenerations = "/v1/images/generations"
EndpointImagesEdits = "/v1/images/edits"
EndpointVideosGenerations = "/v1/videos/generations"
EndpointVideos = "/v1/videos"
EndpointGeminiModels = "/v1beta/models"
)
@@ -53,6 +55,10 @@ func NormalizeInboundEndpoint(path string) string {
return EndpointImagesGenerations
case strings.Contains(path, EndpointImagesEdits) || strings.Contains(path, "/images/edits"):
return EndpointImagesEdits
case strings.Contains(path, EndpointVideosGenerations) || strings.Contains(path, "/videos/generations"):
return EndpointVideosGenerations
case strings.Contains(path, EndpointVideos) || strings.Contains(path, "/videos/"):
return EndpointVideos
case strings.Contains(path, EndpointResponses):
return EndpointResponses
case strings.Contains(path, EndpointGeminiModels):
@@ -78,7 +84,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
switch platform {
case service.PlatformOpenAI, service.PlatformGrok:
if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits {
if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
return inbound
}
// OpenAI forwards everything to the Responses API.
@@ -28,6 +28,8 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/v1/responses", EndpointResponses},
{"/v1/images/generations", EndpointImagesGenerations},
{"/v1/images/edits", EndpointImagesEdits},
{"/v1/videos/generations", EndpointVideosGenerations},
{"/v1/videos/req_123", EndpointVideos},
{"/v1beta/models", EndpointGeminiModels},
// Prefixed paths (antigravity, openai).
@@ -81,6 +83,8 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
{"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings},
{"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations},
{"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits},
{"grok video generations", EndpointVideosGenerations, "/v1/videos/generations", service.PlatformGrok, EndpointVideosGenerations},
{"grok video status", EndpointVideos, "/videos/req_123", service.PlatformGrok, EndpointVideos},
// Antigravity — uses inbound to pick Claude vs Gemini upstream.
{"antigravity claude", EndpointMessages, "/antigravity/v1/messages", service.PlatformAntigravity, EndpointMessages},
+366
View File
@@ -0,0 +1,366 @@
package handler
import (
"context"
"errors"
"net/http"
"strconv"
"strings"
"time"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
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"
)
// GrokImages handles xAI image generation/editing through Grok groups.
func (h *OpenAIGatewayHandler) GrokImages(c *gin.Context) {
endpoint := service.GrokMediaEndpointImagesGenerations
if strings.Contains(c.Request.URL.Path, "/images/edits") {
endpoint = service.GrokMediaEndpointImagesEdits
}
h.handleGrokMedia(c, endpoint, "")
}
// GrokVideoGeneration handles xAI video generation through Grok groups.
func (h *OpenAIGatewayHandler) GrokVideoGeneration(c *gin.Context) {
h.handleGrokMedia(c, service.GrokMediaEndpointVideosGenerations, "")
}
// GrokVideoStatus handles xAI video status retrieval through Grok groups.
func (h *OpenAIGatewayHandler) GrokVideoStatus(c *gin.Context) {
h.handleGrokMedia(c, service.GrokMediaEndpointVideoStatus, c.Param("request_id"))
}
func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.GrokMediaEndpoint, requestID string) {
streamStarted := false
defer h.recoverResponsesPanic(c, &streamStarted)
requestStart := time.Now()
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok {
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
return
}
subject, ok := middleware2.GetAuthSubjectFromContext(c)
if !ok {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
return
}
reqLog := requestLogger(
c,
"handler.openai_gateway.grok_media",
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
zap.String("endpoint", string(endpoint)),
)
if !h.ensureResponsesDependencies(c, reqLog) {
return
}
var body []byte
var err error
if endpoint.RequiresRequestBody() {
body, err = pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
if err != nil {
if maxErr, ok := extractMaxBytesError(err); ok {
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
if len(body) == 0 {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
return
}
}
contentType := c.GetHeader("Content-Type")
requestModel := service.ExtractGrokMediaModel(contentType, body)
if endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) == "" {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
if endpoint == service.GrokMediaEndpointVideoStatus && strings.TrimSpace(requestID) == "" {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "request_id is required")
return
}
reqLog = reqLog.With(zap.String("model", requestModel))
setOpsRequestContext(c, requestModel, false)
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
if endpoint.IsGenerationRequest() {
if !service.GroupAllowsImageGeneration(apiKey.Group) {
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
return
}
if moderationBody := grokMediaModerationBody(body); 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)
return
}
}
imageReleaseFunc, acquired := h.acquireImageGenerationSlot(c, streamStarted)
if !acquired {
return
}
if imageReleaseFunc != nil {
defer imageReleaseFunc()
}
}
if h.errorPassthroughService != nil {
service.BindErrorPassthroughService(c, h.errorPassthroughService)
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
userReleaseFunc, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog)
if !acquired {
return
}
if userReleaseFunc != nil {
defer userReleaseFunc()
}
if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
reqLog.Info("grok_media.billing_eligibility_check_failed", zap.Error(err))
status, code, message, retryAfter := billingErrorDetails(err)
if retryAfter > 0 {
c.Header("Retry-After", strconv.Itoa(retryAfter))
}
h.errorResponse(c, status, code, message)
return
}
sessionSeed := body
if len(sessionSeed) == 0 && strings.TrimSpace(requestID) != "" {
sessionSeed = []byte(requestID)
}
sessionHash := h.gatewayService.GenerateExplicitSessionHash(c, sessionSeed)
requestCtx := c.Request.Context()
failedAccountIDs := make(map[int64]struct{})
sameAccountRetryCount := make(map[int64]int)
var lastFailoverErr *service.UpstreamFailoverError
switchCount := 0
maxAccountSwitches := h.maxAccountSwitches
if maxAccountSwitches <= 0 {
maxAccountSwitches = 3
}
routingStart := time.Now()
for {
selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
requestCtx,
apiKey.GroupID,
"",
sessionHash,
requestModel,
failedAccountIDs,
service.OpenAIUpstreamTransportHTTPSSE,
"",
false,
service.PlatformGrok,
)
if err != nil {
reqLog.Warn("grok_media.account_select_failed",
zap.Error(err),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if len(failedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestModel, requestModel, service.PlatformGrok)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if lastFailoverErr != nil {
h.handleFailoverExhausted(c, lastFailoverErr, false)
} else {
h.errorResponse(c, http.StatusBadGateway, "api_error", "Upstream request failed")
}
return
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestModel, requestModel, service.PlatformGrok)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
reqLog.Debug("grok_media.account_schedule_decision",
zap.String("layer", scheduleDecision.Layer),
zap.Bool("sticky_session_hit", scheduleDecision.StickySessionHit),
zap.Int("candidate_count", scheduleDecision.CandidateCount),
zap.Int("top_k", scheduleDecision.TopK),
zap.Int64("latency_ms", scheduleDecision.LatencyMs),
zap.Float64("load_skew", scheduleDecision.LoadSkew),
)
account := selection.Account
sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account)
setOpsSelectedAccount(c, account.ID, account.Platform)
accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog)
if !accountAcquired {
return
}
service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
forwardStart := time.Now()
writerSizeBeforeForward := c.Writer.Size()
result, err := func() (*service.OpenAIForwardResult, error) {
defer func() {
if accountReleaseFunc != nil {
accountReleaseFunc()
}
}()
return h.gatewayService.ForwardGrokMedia(requestCtx, c, account, endpoint, requestID, body, contentType)
}()
forwardDurationMs := time.Since(forwardStart).Milliseconds()
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
responseLatencyMs := forwardDurationMs
if upstreamLatencyMs > 0 && forwardDurationMs > upstreamLatencyMs {
responseLatencyMs = forwardDurationMs - upstreamLatencyMs
}
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, responseLatencyMs)
if err != nil {
var failoverErr *service.UpstreamFailoverError
if errors.As(err, &failoverErr) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() != writerSizeBeforeForward {
h.handleFailoverExhausted(c, failoverErr, true)
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
sameAccountRetryCount[account.ID]++
reqLog.Warn("grok_media.pool_mode_same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_limit", retryLimit),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
)
select {
case <-requestCtx.Done():
return
case <-time.After(sameAccountRetryDelay):
}
continue
}
}
h.gatewayService.RecordOpenAIAccountSwitch()
failedAccountIDs[account.ID] = struct{}{}
lastFailoverErr = failoverErr
if switchCount >= maxAccountSwitches {
h.handleFailoverExhausted(c, failoverErr, false)
return
}
switchCount++
reqLog.Warn("grok_media.upstream_failover_switching",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("switch_count", switchCount),
zap.Int("max_switches", maxAccountSwitches),
)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
reqLog.Warn("grok_media.forward_failed",
zap.Int64("account_id", account.ID),
zap.Error(err),
)
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, result, requestModel, body, requestID)
reqLog.Debug("grok_media.request_completed",
zap.Int64("account_id", account.ID),
zap.Int("switch_count", switchCount),
)
return
}
}
func grokMediaModerationBody(body []byte) []byte {
if gjson.ValidBytes(body) {
return body
}
return nil
}
func recordGrokMediaUsage(
c *gin.Context,
h *OpenAIGatewayHandler,
reqLog *zap.Logger,
apiKey *service.APIKey,
subject middleware2.AuthSubject,
subscription *service.UserSubscription,
account *service.Account,
result *service.OpenAIForwardResult,
requestModel string,
body []byte,
requestID string,
) {
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
payloadForHash := body
if len(payloadForHash) == 0 && strings.TrimSpace(requestID) != "" {
payloadForHash = []byte(requestID)
}
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
channelUsageFields := service.ChannelUsageFields{
OriginalModel: requestModel,
ChannelMappedModel: requestModel,
}
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
User: apiKey.User,
Account: account,
Subscription: subscription,
InboundEndpoint: inboundEndpoint,
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: service.HashUsageRequestPayload(payloadForHash),
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
ChannelUsageFields: channelUsageFields,
}); err != nil {
logger.L().With(
zap.String("component", "handler.openai_gateway.grok_media"),
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
zap.String("model", requestModel),
zap.Int64("account_id", account.ID),
).Error("grok_media.record_usage_failed", zap.Error(err))
reqLog.Debug("grok_media.record_usage_failed", zap.Error(err))
}
})
}
+36
View File
@@ -437,6 +437,42 @@ func BuildChatCompletionsURL(baseURL string) (string, error) {
return validatedBaseURL + "/chat/completions", nil
}
func BuildImagesGenerationsURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/images/generations", nil
}
func BuildImagesEditsURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/images/edits", nil
}
func BuildVideosGenerationsURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/generations", nil
}
func BuildVideoURL(baseURL, requestID string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
requestID = strings.TrimSpace(requestID)
if requestID == "" {
return "", fmt.Errorf("request id is required")
}
return validatedBaseURL + "/videos/" + url.PathEscape(requestID), nil
}
// TokenResponse represents xAI OAuth token responses.
type TokenResponse struct {
AccessToken string `json:"access_token"`
+21
View File
@@ -116,6 +116,27 @@ func TestValidateXAIURLsAllowOfficialOAuthAndGatewayHosts(t *testing.T) {
require.Equal(t, DefaultCLIBaseURL+"/chat/completions", chatURL)
}
func TestBuildGrokMediaURLs(t *testing.T) {
imagesURL, err := BuildImagesGenerationsURL(DefaultBaseURL + "/")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/images/generations", imagesURL)
editsURL, err := BuildImagesEditsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/images/edits", editsURL)
videosURL, err := BuildVideosGenerationsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/generations", videosURL)
videoURL, err := BuildVideoURL(DefaultBaseURL, "req 123")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/req%20123", videoURL)
_, err = BuildVideoURL(DefaultBaseURL, " ")
require.Error(t, err)
}
func TestValidateXAIURLsRejectArbitraryHostsByDefault(t *testing.T) {
_, err := ValidateOAuthEndpointURL("https://auth.example.test/oauth2/token")
require.Error(t, err)
+50 -52
View File
@@ -42,6 +42,48 @@ func RegisterGatewayRoutes(
isOpenAIGatewayPlatform := func(c *gin.Context) bool {
return getGroupPlatform(c) == service.PlatformOpenAI
}
imagesHandler := func(c *gin.Context) {
switch getGroupPlatform(c) {
case service.PlatformOpenAI:
h.OpenAIGateway.Images(c)
case service.PlatformGrok:
h.OpenAIGateway.GrokImages(c)
default:
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Images API is not supported for this platform",
},
})
}
}
videoGenerationHandler := func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
h.OpenAIGateway.GrokVideoGeneration(c)
return
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Videos API is not supported for this platform",
},
})
}
videoStatusHandler := func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
h.OpenAIGateway.GrokVideoStatus(c)
return
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Videos API is not supported for this platform",
},
})
}
// API网关(Claude API兼容)
gateway := r.Group("/v1")
gateway.Use(bodyLimit)
@@ -120,32 +162,10 @@ func RegisterGatewayRoutes(
}
h.OpenAIGateway.Embeddings(c)
})
gateway.POST("/images/generations", func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Images API is not supported for this platform",
},
})
return
}
h.OpenAIGateway.Images(c)
})
gateway.POST("/images/edits", func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Images API is not supported for this platform",
},
})
return
}
h.OpenAIGateway.Images(c)
})
gateway.POST("/images/generations", imagesHandler)
gateway.POST("/images/edits", imagesHandler)
gateway.POST("/videos/generations", videoGenerationHandler)
gateway.GET("/videos/:request_id", videoStatusHandler)
}
// Gemini 原生 API 兼容层(Gemini SDK/CLI 直连)
@@ -206,32 +226,10 @@ func RegisterGatewayRoutes(
}
h.OpenAIGateway.Embeddings(c)
})
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Images API is not supported for this platform",
},
})
return
}
h.OpenAIGateway.Images(c)
})
r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": "Images API is not supported for this platform",
},
})
return
}
h.OpenAIGateway.Images(c)
})
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoGenerationHandler)
r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoStatusHandler)
// Antigravity 模型列表
r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels)
@@ -83,6 +83,62 @@ func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
}
}
func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformGrok)
for _, path := range []string{
"/v1/images/generations",
"/v1/images/edits",
"/images/generations",
"/images/edits",
"/v1/videos/generations",
"/videos/generations",
} {
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"model":"grok-imagine","prompt":"draw a cat"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should hit Grok media handler", path)
require.NotContains(t, w.Body.String(), "not supported for this platform")
}
for _, path := range []string{
"/v1/videos/request-123",
"/videos/request-123",
} {
req := httptest.NewRequest(http.MethodGet, path, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should hit Grok video handler", path)
require.NotContains(t, w.Body.String(), "not supported for this platform")
}
}
func TestGatewayRoutesNonGrokVideosAreRejectedAtPlatformGate(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformOpenAI)
for _, tc := range []struct {
method string
path string
body string
}{
{http.MethodPost, "/v1/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`},
{http.MethodPost, "/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`},
{http.MethodGet, "/v1/videos/request-123", ""},
{http.MethodGet, "/videos/request-123", ""},
} {
req := httptest.NewRequest(tc.method, tc.path, strings.NewReader(tc.body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code, "method=%s path=%s", tc.method, tc.path)
require.Contains(t, w.Body.String(), "Videos API is not supported for this platform")
}
}
func TestGatewayRoutesGrokAllowsCLICompatibilityEntrypoints(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformGrok)
+220
View File
@@ -0,0 +1,220 @@
package service
import (
"bytes"
"context"
"fmt"
"io"
"mime"
"mime/multipart"
"net/http"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
type GrokMediaEndpoint string
const (
GrokMediaEndpointImagesGenerations GrokMediaEndpoint = "images_generations"
GrokMediaEndpointImagesEdits GrokMediaEndpoint = "images_edits"
GrokMediaEndpointVideosGenerations GrokMediaEndpoint = "videos_generations"
GrokMediaEndpointVideoStatus GrokMediaEndpoint = "video_status"
)
func (e GrokMediaEndpoint) RequiresRequestBody() bool {
return e != GrokMediaEndpointVideoStatus
}
func (e GrokMediaEndpoint) IsGenerationRequest() bool {
switch e {
case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits, GrokMediaEndpointVideosGenerations:
return true
default:
return false
}
}
func (e GrokMediaEndpoint) httpMethod() string {
if e == GrokMediaEndpointVideoStatus {
return http.MethodGet
}
return http.MethodPost
}
func ExtractGrokMediaModel(contentType string, body []byte) string {
if model := strings.TrimSpace(gjson.GetBytes(body, "model").String()); model != "" {
return model
}
return extractGrokMediaMultipartModel(contentType, body)
}
func extractGrokMediaMultipartModel(contentType string, body []byte) string {
mediaType, params, err := mime.ParseMediaType(strings.TrimSpace(contentType))
if err != nil || !strings.EqualFold(mediaType, "multipart/form-data") {
return ""
}
boundary := strings.TrimSpace(params["boundary"])
if boundary == "" {
return ""
}
reader := multipart.NewReader(bytes.NewReader(body), boundary)
for {
part, err := reader.NextPart()
if err == io.EOF {
return ""
}
if err != nil {
return ""
}
if part.FormName() != "model" || part.FileName() != "" {
continue
}
data, err := io.ReadAll(part)
if err != nil {
return ""
}
return strings.TrimSpace(string(data))
}
}
func (e GrokMediaEndpoint) upstreamURL(baseURL, requestID string) (string, error) {
switch e {
case GrokMediaEndpointImagesGenerations:
return xai.BuildImagesGenerationsURL(baseURL)
case GrokMediaEndpointImagesEdits:
return xai.BuildImagesEditsURL(baseURL)
case GrokMediaEndpointVideosGenerations:
return xai.BuildVideosGenerationsURL(baseURL)
case GrokMediaEndpointVideoStatus:
return xai.BuildVideoURL(baseURL, requestID)
default:
return "", fmt.Errorf("unsupported grok media endpoint: %s", e)
}
}
func (s *OpenAIGatewayService) ForwardGrokMedia(
ctx context.Context,
c *gin.Context,
account *Account,
endpoint GrokMediaEndpoint,
requestID string,
body []byte,
contentType string,
) (*OpenAIForwardResult, error) {
startTime := time.Now()
if account == nil {
return nil, fmt.Errorf("grok account is required")
}
if account.Platform != PlatformGrok {
return nil, fmt.Errorf("account platform %s is not supported for grok media", account.Platform)
}
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
return nil, err
}
targetURL, err := endpoint.upstreamURL(account.GetGrokBaseURL(), requestID)
if err != nil {
return nil, err
}
var bodyReader io.Reader
if endpoint.RequiresRequestBody() {
bodyReader = bytes.NewReader(body)
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
defer releaseUpstreamCtx()
upstreamReq, err := http.NewRequestWithContext(upstreamCtx, endpoint.httpMethod(), targetURL, bodyReader)
if err != nil {
return nil, err
}
upstreamReq.Header.Set("Authorization", "Bearer "+token)
upstreamReq.Header.Set("Accept", "application/json")
upstreamReq.Header.Set("User-Agent", "sub2api-grok/1.0")
if endpoint.RequiresRequestBody() {
contentType = strings.TrimSpace(contentType)
if contentType == "" {
contentType = "application/json"
}
upstreamReq.Header.Set("Content-Type", contentType)
}
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
upstreamStart := time.Now()
resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds())
if err != nil {
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false)
}
defer func() { _ = resp.Body.Close() }()
requestIDHeader := firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id"))
requestModel := ExtractGrokMediaModel(contentType, body)
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
}
s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode))
respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
if err != nil {
return nil, err
}
writeGrokMediaResponse(c, resp, respBody, s.responseHeaderFilter)
return &OpenAIForwardResult{
RequestID: requestIDHeader,
Model: requestModel,
UpstreamModel: requestModel,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
}, nil
}
func writeGrokMediaResponse(c *gin.Context, resp *http.Response, body []byte, filter *responseheaders.CompiledHeaderFilter) {
if c == nil || resp == nil {
return
}
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, filter)
contentType := strings.TrimSpace(resp.Header.Get("Content-Type"))
if contentType == "" {
contentType = "application/json"
}
c.Data(resp.StatusCode, contentType, body)
}
@@ -7,6 +7,7 @@ import (
"context"
"encoding/json"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"strings"
@@ -144,6 +145,103 @@ func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) {
require.Contains(t, err.Error(), "invalid base url")
}
func TestExtractGrokMediaModelSupportsJSONAndMultipart(t *testing.T) {
require.Equal(t, "grok-imagine", ExtractGrokMediaModel("application/json", []byte(`{"model":"grok-imagine"}`)))
var buf bytes.Buffer
writer := multipart.NewWriter(&buf)
require.NoError(t, writer.WriteField("prompt", "draw a cat"))
require.NoError(t, writer.WriteField("model", "grok-imagine-edit"))
require.NoError(t, writer.Close())
require.Equal(t, "grok-imagine-edit", ExtractGrokMediaModel(writer.FormDataContentType(), buf.Bytes()))
}
func TestForwardGrokMediaImagesGenerationPassthrough(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: 61,
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-image-req"},
},
Body: io.NopCloser(strings.NewReader(`{"data":[]}`)),
}}
svc := &OpenAIGatewayService{httpUpstream: upstream}
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesGenerations, "", body, "application/json")
require.NoError(t, err)
require.Equal(t, "https://xai.test/v1/images/generations", upstream.lastReq.URL.String())
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.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)
}
func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/videos/request-123", nil)
account := &Account{
ID: 62,
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-req"},
},
Body: io.NopCloser(strings.NewReader(`{"id":"request-123","status":"completed"}`)),
}}
svc := &OpenAIGatewayService{httpUpstream: upstream}
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideoStatus, "request-123", nil, "")
require.NoError(t, err)
require.Equal(t, "https://xai.test/v1/videos/request-123", upstream.lastReq.URL.String())
require.Equal(t, http.MethodGet, upstream.lastReq.Method)
require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization"))
require.Empty(t, upstream.lastReq.Header.Get("Content-Type"))
require.Empty(t, upstream.lastBody)
require.Equal(t, http.StatusOK, recorder.Code)
require.JSONEq(t, `{"id":"request-123","status":"completed"}`, recorder.Body.String())
require.Equal(t, "xai-video-req", result.RequestID)
}
func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *testing.T) {
gin.SetMode(gin.TestMode)