feat: support Grok video edits and extensions

This commit is contained in:
chenjian
2026-07-13 10:50:42 +08:00
parent b73d8c3efe
commit 909b96edd2
11 changed files with 135 additions and 6 deletions
+7 -1
View File
@@ -24,6 +24,8 @@ const (
EndpointImagesGenerations = "/v1/images/generations"
EndpointImagesEdits = "/v1/images/edits"
EndpointVideosGenerations = "/v1/videos/generations"
EndpointVideosEdits = "/v1/videos/edits"
EndpointVideosExtensions = "/v1/videos/extensions"
EndpointVideos = "/v1/videos"
EndpointGeminiModels = "/v1beta/models"
)
@@ -88,6 +90,10 @@ func NormalizeInboundEndpoint(path string) string {
return EndpointImagesEdits
case strings.Contains(path, EndpointVideosGenerations) || strings.Contains(path, "/videos/generations"):
return EndpointVideosGenerations
case strings.Contains(path, EndpointVideosEdits) || strings.Contains(path, "/videos/edits"):
return EndpointVideosEdits
case strings.Contains(path, EndpointVideosExtensions) || strings.Contains(path, "/videos/extensions"):
return EndpointVideosExtensions
case strings.Contains(path, EndpointVideos) || strings.Contains(path, "/videos/"):
return EndpointVideos
case strings.Contains(path, EndpointResponsesCompact) || isResponsesCompactAliasPath(path):
@@ -173,7 +179,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
switch platform {
case service.PlatformOpenAI, service.PlatformGrok:
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideosEdits || inbound == EndpointVideosExtensions || inbound == EndpointVideos {
return inbound
}
// OpenAI forwards everything to the Responses API.
+11 -1
View File
@@ -31,6 +31,16 @@ func (h *OpenAIGatewayHandler) GrokVideoGeneration(c *gin.Context) {
h.handleGrokMedia(c, service.GrokMediaEndpointVideosGenerations, "")
}
// GrokVideoEdit handles asynchronous xAI video edits through Grok groups.
func (h *OpenAIGatewayHandler) GrokVideoEdit(c *gin.Context) {
h.handleGrokMedia(c, service.GrokMediaEndpointVideosEdits, "")
}
// GrokVideoExtension handles asynchronous xAI video extensions through Grok groups.
func (h *OpenAIGatewayHandler) GrokVideoExtension(c *gin.Context) {
h.handleGrokMedia(c, service.GrokMediaEndpointVideosExtensions, "")
}
// 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"))
@@ -298,7 +308,7 @@ 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 endpoint.IsGenerationRequest() && 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),
+16
View File
@@ -461,6 +461,22 @@ func BuildVideosGenerationsURL(baseURL string) (string, error) {
return validatedBaseURL + "/videos/generations", nil
}
func BuildVideosEditsURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/edits", nil
}
func BuildVideosExtensionsURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/extensions", nil
}
func BuildVideoURL(baseURL, requestID string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
+8
View File
@@ -129,6 +129,14 @@ func TestBuildGrokMediaURLs(t *testing.T) {
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/generations", videosURL)
videoEditsURL, err := BuildVideosEditsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/edits", videoEditsURL)
videoExtensionsURL, err := BuildVideosExtensionsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/extensions", videoExtensionsURL)
videoURL, err := BuildVideoURL(DefaultBaseURL, "req 123")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/req%20123", videoURL)
+20
View File
@@ -84,6 +84,22 @@ func RegisterGatewayRoutes(
},
})
}
videoEditHandler := func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
h.OpenAIGateway.GrokVideoEdit(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"}})
}
videoExtensionHandler := func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
h.OpenAIGateway.GrokVideoExtension(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)
@@ -185,6 +201,8 @@ func RegisterGatewayRoutes(
gateway.DELETE("/images/batches/:id", h.BatchImage.DeleteRecord)
gateway.DELETE("/images/batches/:id/outputs", h.BatchImage.DeleteOutputs)
gateway.POST("/videos/generations", videoGenerationHandler)
gateway.POST("/videos/edits", videoEditHandler)
gateway.POST("/videos/extensions", videoExtensionHandler)
gateway.GET("/videos/:request_id", videoStatusHandler)
}
@@ -252,6 +270,8 @@ func RegisterGatewayRoutes(
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.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoEditHandler)
r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoExtensionHandler)
r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoStatusHandler)
// Antigravity 模型列表
@@ -123,6 +123,10 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) {
"/images/edits",
"/v1/videos/generations",
"/videos/generations",
"/v1/videos/edits",
"/videos/edits",
"/v1/videos/extensions",
"/videos/extensions",
} {
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"model":"grok-imagine","prompt":"draw a cat"}`))
req.Header.Set("Content-Type", "application/json")
@@ -156,6 +160,10 @@ func TestGatewayRoutesNonGrokVideosAreRejectedAtPlatformGate(t *testing.T) {
}{
{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.MethodPost, "/v1/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`},
{http.MethodPost, "/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`},
{http.MethodPost, "/v1/videos/extensions", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`},
{http.MethodPost, "/videos/extensions", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`},
{http.MethodGet, "/v1/videos/request-123", ""},
{http.MethodGet, "/videos/request-123", ""},
} {
+8 -2
View File
@@ -26,6 +26,8 @@ const (
GrokMediaEndpointImagesGenerations GrokMediaEndpoint = "images_generations"
GrokMediaEndpointImagesEdits GrokMediaEndpoint = "images_edits"
GrokMediaEndpointVideosGenerations GrokMediaEndpoint = "videos_generations"
GrokMediaEndpointVideosEdits GrokMediaEndpoint = "videos_edits"
GrokMediaEndpointVideosExtensions GrokMediaEndpoint = "videos_extensions"
GrokMediaEndpointVideoStatus GrokMediaEndpoint = "video_status"
)
@@ -35,7 +37,7 @@ func (e GrokMediaEndpoint) RequiresRequestBody() bool {
func (e GrokMediaEndpoint) IsGenerationRequest() bool {
switch e {
case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits, GrokMediaEndpointVideosGenerations:
case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits, GrokMediaEndpointVideosGenerations, GrokMediaEndpointVideosEdits, GrokMediaEndpointVideosExtensions:
return true
default:
return false
@@ -274,6 +276,10 @@ func (e GrokMediaEndpoint) upstreamURL(baseURL, requestID string) (string, error
return xai.BuildImagesEditsURL(baseURL)
case GrokMediaEndpointVideosGenerations:
return xai.BuildVideosGenerationsURL(baseURL)
case GrokMediaEndpointVideosEdits:
return xai.BuildVideosEditsURL(baseURL)
case GrokMediaEndpointVideosExtensions:
return xai.BuildVideosExtensionsURL(baseURL)
case GrokMediaEndpointVideoStatus:
return xai.BuildVideoURL(baseURL, requestID)
default:
@@ -531,7 +537,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi
meta.ImageSize = requestInfo.SizeTier
meta.ImageInputSize = requestInfo.Size
meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody)
case GrokMediaEndpointVideosGenerations:
case GrokMediaEndpointVideosGenerations, GrokMediaEndpointVideosEdits, GrokMediaEndpointVideosExtensions:
meta.ResponseID = extractGrokMediaVideoRequestID(responseBody)
meta.VideoCount = 1
meta.VideoResolution = requestInfo.Resolution
@@ -637,6 +637,49 @@ func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) {
require.Equal(t, "xai-video-req", result.RequestID)
}
func TestForwardGrokMediaVideoMutationEndpoints(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
tests := []struct {
name string
endpoint GrokMediaEndpoint
path string
}{
{name: "edit", endpoint: GrokMediaEndpointVideosEdits, path: "/videos/edits"},
{name: "extension", endpoint: GrokMediaEndpointVideosExtensions, path: "/videos/extensions"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok-imagine-video","prompt":"continue","video":{"url":"https://example.com/in.mp4"},"duration":6}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1"+tt.path, bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
account := &Account{
ID: 71, 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"}},
Body: io.NopCloser(strings.NewReader(`{"request_id":"video-mutation-123"}`)),
}}
svc := &OpenAIGatewayService{httpUpstream: upstream}
result, err := svc.ForwardGrokMedia(context.Background(), c, account, tt.endpoint, "", body, "application/json")
require.NoError(t, err)
require.Equal(t, "https://xai.test/v1"+tt.path, upstream.lastReq.URL.String())
require.Equal(t, http.MethodPost, upstream.lastReq.Method)
require.JSONEq(t, string(body), string(upstream.lastBody))
require.Equal(t, "video-mutation-123", result.ResponseID)
require.Equal(t, 1, result.VideoCount)
require.Equal(t, 6, result.VideoDurationSeconds)
})
}
}
func TestBindGrokMediaVideoRequestAccountUsesRequestIDStickyHash(t *testing.T) {
ctx := context.Background()
groupID := int64(7)
+2 -1
View File
@@ -308,7 +308,8 @@ func shouldBypassEmbeddedFrontend(path string) bool {
trimmed == "/health" ||
trimmed == "/responses" ||
strings.HasPrefix(trimmed, "/responses/") ||
strings.HasPrefix(trimmed, "/images/")
strings.HasPrefix(trimmed, "/images/") ||
strings.HasPrefix(trimmed, "/videos/")
}
func serveIndexHTML(c *gin.Context, fsys fs.FS) {
+11
View File
@@ -562,6 +562,17 @@ func TestFrontendServer_Middleware(t *testing.T) {
})
}
func TestEmbeddedFrontendBypassesBareVideoAPIRoutes(t *testing.T) {
for _, path := range []string{
"/videos/generations",
"/videos/edits",
"/videos/extensions",
"/videos/request-123",
} {
require.True(t, shouldBypassEmbeddedFrontend(path), "path=%s", path)
}
}
func TestNewFrontendServer(t *testing.T) {
t.Run("creates_server_successfully", func(t *testing.T) {
provider := &mockSettingsProvider{