diff --git a/sdk/api/handlers/openai/openai_videos_handlers.go b/sdk/api/handlers/openai/openai_videos_handlers.go index 76573ccf..6f9e1c3a 100644 --- a/sdk/api/handlers/openai/openai_videos_handlers.go +++ b/sdk/api/handlers/openai/openai_videos_handlers.go @@ -46,12 +46,12 @@ const defaultVideoAuthBindingTTL = 3 * time.Hour var videoAuthBindings = newVideoAuthBindingStore() type xaiVideoCreateMetadata struct { - Model string - UpstreamModel string - Prompt string - Seconds string - Size string - CreatedAt int64 + Model string + RoutingModel string + Prompt string + Seconds string + Size string + CreatedAt int64 } type videoAuthBinding struct { @@ -206,6 +206,21 @@ func canonicalXAIVideosModel(model string) string { return defaultXAIVideosModel } +func routingXAIVideosModel(model string) string { + if isSoraVideosModel(model) { + return defaultXAIVideosModel + } + switch videosModelBase(model) { + case defaultXAIVideosModel: + return defaultXAIVideosModel + case xaiVideos15Model: + return xaiVideos15Model + case xaiVideos15PreviewAlias: + return xaiVideos15PreviewAlias + } + return defaultXAIVideosModel +} + func responseVideosModel(model string) string { return canonicalXAIVideosModel(model) } @@ -299,11 +314,11 @@ func (h *OpenAIAPIHandler) bindVideoAuthIDAndModelFromPayload(payload []byte, au if videoID == "" { return } - videoAuthBindings.setWithModel(videoID, authID, canonicalXAIVideosModel(model), h.videoAuthBindingTTL()) + videoAuthBindings.setWithModel(videoID, authID, routingXAIVideosModel(model), h.videoAuthBindingTTL()) } func (h *OpenAIAPIHandler) bindVideoAuthID(videoID string, authID string, model string) { - videoAuthBindings.setWithModel(videoID, authID, canonicalXAIVideosModel(model), h.videoAuthBindingTTL()) + videoAuthBindings.setWithModel(videoID, authID, routingXAIVideosModel(model), h.videoAuthBindingTTL()) } func (h *OpenAIAPIHandler) contextWithVideoAuthBinding(ctx context.Context, videoID string) context.Context { @@ -375,12 +390,12 @@ func buildXAIVideosCreateRequest(rawJSON []byte, model string) ([]byte, xaiVideo } meta := xaiVideoCreateMetadata{ - Model: responseVideosModel(model), - UpstreamModel: videoModel, - Prompt: prompt, - Seconds: seconds, - Size: size, - CreatedAt: time.Now().Unix(), + Model: responseVideosModel(model), + RoutingModel: routingXAIVideosModel(model), + Prompt: prompt, + Seconds: seconds, + Size: size, + CreatedAt: time.Now().Unix(), } return req, meta, nil } @@ -733,9 +748,9 @@ func (h *OpenAIAPIHandler) handleXAIVideosNativePost(c *gin.Context) { return } - videoModel = canonicalXAIVideosModel(videoModel) - rawJSON, _ = sjson.SetBytes(rawJSON, "model", videoModel) - h.collectXAIVideosNative(c, rawJSON, videoModel, true) + routingModel := routingXAIVideosModel(videoModel) + rawJSON, _ = sjson.SetBytes(rawJSON, "model", canonicalXAIVideosModel(videoModel)) + h.collectXAIVideosNative(c, rawJSON, routingModel, true) } func (h *OpenAIAPIHandler) XAIVideosRetrieve(c *gin.Context) { @@ -998,12 +1013,12 @@ func (h *OpenAIAPIHandler) collectXAIVideosCreate(c *gin.Context, xaiReq []byte, cliCtx = handlers.WithSelectedAuthIDCallback(cliCtx, func(authID string) { selectedAuthID = authID }) - upstreamModel := strings.TrimSpace(meta.UpstreamModel) - if upstreamModel == "" { - upstreamModel = meta.Model + routingModel := strings.TrimSpace(meta.RoutingModel) + if routingModel == "" { + routingModel = routingXAIVideosModel(meta.Model) } stopKeepAlive := h.StartNonStreamingKeepAlive(c, cliCtx) - resp, upstreamHeaders, errMsg := h.ExecuteWithAuthManager(cliCtx, xaiVideosHandlerType, upstreamModel, xaiReq, "") + resp, upstreamHeaders, errMsg := h.ExecuteWithAuthManager(cliCtx, xaiVideosHandlerType, routingModel, xaiReq, "") stopKeepAlive() if errMsg != nil { h.WriteErrorResponse(c, errMsg) @@ -1023,7 +1038,7 @@ func (h *OpenAIAPIHandler) collectXAIVideosCreate(c *gin.Context, xaiReq []byte, return } - h.bindVideoAuthIDFromPayload(out, selectedAuthID) + h.bindVideoAuthIDAndModelFromPayload(out, selectedAuthID, routingModel) handlers.WriteUpstreamHeaders(c.Writer.Header(), upstreamHeaders) _, _ = c.Writer.Write(out) cliCancel(nil) diff --git a/sdk/api/handlers/openai/openai_videos_handlers_test.go b/sdk/api/handlers/openai/openai_videos_handlers_test.go index c4a82c29..8b2a7afa 100644 --- a/sdk/api/handlers/openai/openai_videos_handlers_test.go +++ b/sdk/api/handlers/openai/openai_videos_handlers_test.go @@ -271,6 +271,9 @@ func TestBuildXAIVideosCreateRequestAllowsVideo15Model(t *testing.T) { if meta.Model != xaiVideos15Model { t.Fatalf("meta model = %q, want %s", meta.Model, xaiVideos15Model) } + if meta.RoutingModel != xaiVideos15Model { + t.Fatalf("routing model = %q, want %s", meta.RoutingModel, xaiVideos15Model) + } } func TestBuildXAIVideosCreateRequestNormalizesVideo15PreviewAlias(t *testing.T) { @@ -287,6 +290,9 @@ func TestBuildXAIVideosCreateRequestNormalizesVideo15PreviewAlias(t *testing.T) if meta.Model != xaiVideos15Model { t.Fatalf("meta model = %q, want %s", meta.Model, xaiVideos15Model) } + if meta.RoutingModel != xaiVideos15PreviewAlias { + t.Fatalf("routing model = %q, want %s", meta.RoutingModel, xaiVideos15PreviewAlias) + } } func TestBuildXAIVideosCreateRequestAllowsCustomSeconds(t *testing.T) { @@ -842,6 +848,94 @@ func TestXAIVideosNativeRetrieveUsesCanonicalBoundModel(t *testing.T) { } } +func TestVideosCreatePreviewAliasUsesPreviewAuthWithGAPayload(t *testing.T) { + resetVideoAuthBindingsForTest(t) + executor := &videoAuthCaptureExecutor{requestID: "video-openai-preview-alias"} + handler := newVideoSingleModelAuthTestHandler(t, executor, "video-openai-preview-auth", xaiVideos15PreviewAlias) + + createResp := performVideosEndpointRequest(t, http.MethodPost, openAIVideosPath, "application/json", strings.NewReader(`{"model":"grok-imagine-video-1.5-preview","prompt":"make a video"}`), handler.VideosCreate) + if createResp.Code != http.StatusOK { + t.Fatalf("create status = %d, want %d: %s", createResp.Code, http.StatusOK, createResp.Body.String()) + } + videoID := gjson.GetBytes(createResp.Body.Bytes(), "id").String() + if got := gjson.GetBytes(createResp.Body.Bytes(), "model").String(); got != xaiVideos15Model { + t.Fatalf("response model = %q, want %s", got, xaiVideos15Model) + } + + retrieveResp := performVideosRouteRequest(t, http.MethodGet, openAIVideosPath+"/:video_id", openAIVideosPath+"/"+videoID, "", nil, handler.VideosRetrieve) + if retrieveResp.Code != http.StatusOK { + t.Fatalf("retrieve status = %d, want %d: %s", retrieveResp.Code, http.StatusOK, retrieveResp.Body.String()) + } + + assertPreviewAliasRouting(t, executor, videoID, "video-openai-preview-auth") +} + +func TestXAIVideosNativePreviewAliasUsesPreviewAuthWithGAPayload(t *testing.T) { + resetVideoAuthBindingsForTest(t) + executor := &videoAuthCaptureExecutor{requestID: "video-native-preview-alias"} + handler := newVideoSingleModelAuthTestHandler(t, executor, "video-native-preview-auth", xaiVideos15PreviewAlias) + + createResp := performVideosEndpointRequest(t, http.MethodPost, xaiVideosGenerationsAPI, "application/json", strings.NewReader(`{"model":"grok-imagine-video-1.5-preview","prompt":"make a video"}`), handler.XAIVideosGenerations) + if createResp.Code != http.StatusOK { + t.Fatalf("create status = %d, want %d: %s", createResp.Code, http.StatusOK, createResp.Body.String()) + } + videoID := gjson.GetBytes(createResp.Body.Bytes(), "request_id").String() + + retrieveResp := performVideosRouteRequest(t, http.MethodGet, videosPath+"/:request_id", videosPath+"/"+videoID, "", nil, handler.XAIVideosRetrieve) + if retrieveResp.Code != http.StatusOK { + t.Fatalf("retrieve status = %d, want %d: %s", retrieveResp.Code, http.StatusOK, retrieveResp.Body.String()) + } + + assertPreviewAliasRouting(t, executor, videoID, "video-native-preview-auth") +} + +func newVideoSingleModelAuthTestHandler(t *testing.T, executor *videoAuthCaptureExecutor, authID string, model string) *OpenAIAPIHandler { + t.Helper() + + manager := coreauth.NewManager(nil, &coreauth.RoundRobinSelector{}, nil) + manager.RegisterExecutor(executor) + auth := &coreauth.Auth{ + ID: authID, + Provider: "xai", + Status: coreauth.StatusActive, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("manager.Register(%s): %v", authID, errRegister) + } + registry.GetGlobalRegistry().RegisterClient(authID, auth.Provider, []*registry.ModelInfo{{ID: model}}) + manager.RefreshSchedulerEntry(authID) + t.Cleanup(func() { + registry.GetGlobalRegistry().UnregisterClient(authID) + }) + + base := apihandlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, manager) + return NewOpenAIAPIHandler(base) +} + +func assertPreviewAliasRouting(t *testing.T, executor *videoAuthCaptureExecutor, videoID string, authID string) { + t.Helper() + + authIDs := executor.AuthIDs() + if len(authIDs) != 2 || authIDs[0] != authID || authIDs[1] != authID { + t.Fatalf("authIDs = %v, want both calls to use %s", authIDs, authID) + } + models := executor.Models() + if len(models) != 2 || models[0] != xaiVideos15PreviewAlias || models[1] != xaiVideos15PreviewAlias { + t.Fatalf("models = %v, want both calls to route with %s", models, xaiVideos15PreviewAlias) + } + payloadModels := executor.PayloadModels() + if len(payloadModels) != 2 || payloadModels[0] != xaiVideos15Model { + t.Fatalf("payload models = %v, want create payload model %s", payloadModels, xaiVideos15Model) + } + binding, ok := videoAuthBindings.getBinding(videoID) + if !ok { + t.Fatal("video auth binding was not stored") + } + if binding.authID != authID || binding.model != xaiVideos15PreviewAlias { + t.Fatalf("binding = {authID:%q model:%q}, want {authID:%q model:%q}", binding.authID, binding.model, authID, xaiVideos15PreviewAlias) + } +} + func TestVideoAuthBindingTTLUsesConfig(t *testing.T) { base := apihandlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{VideoResultAuthCacheTTL: "45m"}, nil) handler := NewOpenAIAPIHandler(base)