mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat: support Grok video edits and extensions
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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", ""},
|
||||
} {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user