From 72bbf8b87e00b2c0e3f3755a180dd25e934c0181 Mon Sep 17 00:00:00 2001 From: Chenyme <118253778+chenyme@users.noreply.github.com> Date: Fri, 7 Aug 2026 10:12:26 +0800 Subject: [PATCH 1/3] feat: add native xAI image and video generation support --- .../application/channel/model_catalog.go | 14 +- .../application/channel/model_catalog_test.go | 11 +- .../conversation/model_option_policy.go | 4 + .../conversation/model_option_policy_test.go | 43 +++ .../conversation/service_media_video.go | 5 +- .../settings/model_option_policy.go | 1 + .../internal/application/settings/service.go | 12 +- .../application/settings/service_seed_test.go | 27 ++ backend/internal/infra/config/config.go | 5 + backend/internal/infra/llm/adapter.go | 15 +- backend/internal/infra/llm/adapter_test.go | 15 + backend/internal/infra/llm/client.go | 5 + .../internal/infra/llm/endpoint_url_test.go | 6 + backend/internal/infra/llm/openai.go | 2 + backend/internal/infra/llm/xai_images.go | 108 ++++-- backend/internal/infra/llm/xai_images_test.go | 44 ++- backend/internal/infra/llm/xai_videos.go | 361 ++++++++++++++++++ backend/internal/infra/llm/xai_videos_test.go | 197 ++++++++++ .../internal/infra/mediaartifact/client.go | 44 ++- .../infra/mediaartifact/client_test.go | 52 ++- frontend/features/admin/api/llm.types.ts | 3 +- .../conversation/admin-conversation.tsx | 7 +- .../models/models-capabilities-presets.tsx | 104 +++++ .../sections/upstreams/upstreams-sheet.tsx | 1 + .../admin/model/conversation-settings.ts | 5 + frontend/features/admin/types/llm.ts | 1 + frontend/features/admin/utils/llm-display.ts | 1 + .../components/sections/chat-model-config.tsx | 4 + frontend/i18n/messages/en-US/chat.json | 4 +- frontend/i18n/messages/zh-CN/chat.json | 4 +- frontend/shared/lib/model-option-policy.ts | 4 + 31 files changed, 1040 insertions(+), 69 deletions(-) create mode 100644 backend/internal/infra/llm/xai_videos.go create mode 100644 backend/internal/infra/llm/xai_videos_test.go diff --git a/backend/internal/application/channel/model_catalog.go b/backend/internal/application/channel/model_catalog.go index 68c74b82b..e19bba7a9 100644 --- a/backend/internal/application/channel/model_catalog.go +++ b/backend/internal/application/channel/model_catalog.go @@ -37,6 +37,7 @@ const ( protocolGeminiInteractions = llm.AdapterGeminiInteractions protocolXAIImage = llm.AdapterXAIImage protocolXAIImageEdits = llm.AdapterXAIImageEdits + protocolXAIVideo = llm.AdapterXAIVideo ) var protocolDefaultKindOrder = []string{ @@ -137,6 +138,7 @@ func systemFallbackProtocols(compatible string) map[string]string { modelKindAudio: llm.AdapterXAIResponses, modelKindImageGen: protocolXAIImage, modelKindImageEdit: protocolXAIImageEdits, + modelKindVideoGen: protocolXAIVideo, } case compatibleOpenRouter: return map[string]string{ @@ -173,7 +175,8 @@ func isKnownProtocol(raw string) bool { protocolGoogleImageGeneration, protocolGeminiInteractions, protocolXAIImage, - protocolXAIImageEdits: + protocolXAIImageEdits, + protocolXAIVideo: return true default: return false @@ -412,7 +415,8 @@ func isProtocolAllowedForKind(kind string, protocol string) bool { case modelKindVideoGen: switch protocol { case protocolOpenAIVideoGenerations, - protocolGeminiInteractions: + protocolGeminiInteractions, + protocolXAIVideo: return true default: return false @@ -504,7 +508,7 @@ func inferKindsJSON(platformModelName string) string { case code == "dall-e-3", strings.HasPrefix(code, "imagen-"): return `["image_gen"]` case code == "sora", code == "veo-2", strings.HasPrefix(code, "kling"), - strings.HasPrefix(code, "veo-"): + strings.HasPrefix(code, "veo-"), isXAIVideoGenerationModel(code): return `["video_gen"]` case strings.HasPrefix(code, "gpt-4o-audio"): return `["audio"]` @@ -539,3 +543,7 @@ func isGeminiImageGenerationModel(code string) bool { func isXAIImageGenerationModel(code string) bool { return strings.HasPrefix(strings.TrimSpace(strings.ToLower(code)), "grok-imagine-image") } + +func isXAIVideoGenerationModel(code string) bool { + return strings.HasPrefix(strings.TrimSpace(strings.ToLower(code)), "grok-imagine-video") +} diff --git a/backend/internal/application/channel/model_catalog_test.go b/backend/internal/application/channel/model_catalog_test.go index 47d77955a..183e45b26 100644 --- a/backend/internal/application/channel/model_catalog_test.go +++ b/backend/internal/application/channel/model_catalog_test.go @@ -51,10 +51,8 @@ func TestProtocolDefaultsForXAIUsesXAIResponsesForConversationKinds(t *testing.T if defaults[modelKindImageEdit] != "xai_image_edits" { t.Fatalf("expected xAI image edit default, got %q in %s", defaults[modelKindImageEdit], raw) } - for _, kind := range []string{modelKindVideoGen} { - if _, ok := defaults[kind]; ok { - t.Fatalf("unexpected xAI default protocol for %s in %s", kind, raw) - } + if defaults[modelKindVideoGen] != "xai_video" { + t.Fatalf("expected xAI video default, got %q in %s", defaults[modelKindVideoGen], raw) } } @@ -448,7 +446,7 @@ func TestInferKindsJSONRecognizesGeminiOmniInteractionsModel(t *testing.T) { } func TestInferKindsJSONRecognizesVideoOnlyModels(t *testing.T) { - for _, modelName := range []string{"veo-3.1-fast"} { + for _, modelName := range []string{"veo-3.1-fast", "grok-imagine-video", "grok-imagine-video-1.5-preview"} { if got := inferKindsJSON(modelName); got != `["video_gen"]` { t.Fatalf("expected %s to infer video generation kind, got %s", modelName, got) } @@ -698,6 +696,9 @@ func TestIsRouteAllowedForTaskSeparatesChatAndImageProtocols(t *testing.T) { if !IsRouteAllowedForTask(TaskTypeVideoGeneration, `["video_gen"]`, "openai_video_generations") { t.Fatalf("expected video generation task to allow OpenAI video protocol") } + if !IsRouteAllowedForTask(TaskTypeVideoGeneration, `["video_gen"]`, "xai_video") { + t.Fatalf("expected video generation task to allow xAI video protocol") + } if IsRouteAllowedForTask(TaskTypeVideoGeneration, `["chat"]`, "openai_responses") { t.Fatalf("expected video generation task to reject chat protocol") } diff --git a/backend/internal/application/conversation/model_option_policy.go b/backend/internal/application/conversation/model_option_policy.go index 96b8d459f..b158c6901 100644 --- a/backend/internal/application/conversation/model_option_policy.go +++ b/backend/internal/application/conversation/model_option_policy.go @@ -497,6 +497,8 @@ func sanitizeModelOptionValues(options map[string]interface{}, protocolKey strin switch protocolKey { case "openai_chat_completions", "openai_responses", "openrouter_responses": sanitizeOpenAIServiceTier(options) + case "xai_video": + llm.SanitizeXAIVideoOptions(options) case "openai_image_generations", "openai_image_edits": value, ok := modelParamIntFromOption(options["partial_images"]) if !ok { @@ -574,6 +576,8 @@ func modelOptionPolicyProtocolKey(protocol string) string { return "xai_image" case llm.AdapterXAIImageEdits: return "xai_image_edits" + case llm.AdapterXAIVideo: + return "xai_video" case llm.AdapterXAIResponses: return "xai_responses" default: diff --git a/backend/internal/application/conversation/model_option_policy_test.go b/backend/internal/application/conversation/model_option_policy_test.go index 99d76816f..ef8fb7b74 100644 --- a/backend/internal/application/conversation/model_option_policy_test.go +++ b/backend/internal/application/conversation/model_option_policy_test.go @@ -1062,6 +1062,49 @@ func TestFilterModelOptionsXAIImageAllowsImageParams(t *testing.T) { } } +func TestFilterModelOptionsXAIVideoAllowsVideoParams(t *testing.T) { + filtered := filterModelOptions(map[string]interface{}{ + "aspect_ratio": " 16:9 ", + "duration": float64(8), + "resolution": "720P", + "prompt": "override", + "image": map[string]interface{}{"url": "https://example.com/source.png"}, + "output": "must not pass through", + }, llm.AdapterXAIVideo, modelOptionPolicyConfig{ + Mode: modelOptionPolicyAllowlist, + AllowedPathsJSON: config.DefaultModelOptionAllowedPathsJSON(), + DeniedPathsJSON: config.DefaultModelOptionDeniedPathsJSON(), + }) + + if filtered["aspect_ratio"] != "16:9" || filtered["duration"] != 8 || filtered["resolution"] != "720p" { + t.Fatalf("expected xAI video params to pass, got %#v", filtered) + } + for _, key := range []string{"prompt", "image", "output"} { + if _, ok := filtered[key]; ok { + t.Fatalf("expected %s to be removed, got %#v", key, filtered) + } + } +} + +func TestFilterModelOptionsXAIVideoDropsInvalidBillableParams(t *testing.T) { + filtered := filterModelOptions(map[string]interface{}{ + "aspect_ratio": "21:9", + "duration": 999, + "resolution": "4k", + }, llm.AdapterXAIVideo, modelOptionPolicyConfig{ + Mode: modelOptionPolicyAllowlist, + AllowedPathsJSON: config.DefaultModelOptionAllowedPathsJSON(), + DeniedPathsJSON: config.DefaultModelOptionDeniedPathsJSON(), + }) + + if len(filtered) != 0 { + t.Fatalf("expected invalid xAI video params to be removed, got %#v", filtered) + } + if duration := mediaDurationSecondsFromOptions(filtered); duration != 0 { + t.Fatalf("expected removed duration not to affect billing, got %d", duration) + } +} + func TestPromptCarriesAssistantReasoning(t *testing.T) { cases := map[string]struct { messages []llm.Message diff --git a/backend/internal/application/conversation/service_media_video.go b/backend/internal/application/conversation/service_media_video.go index 5ff776432..360f6f414 100644 --- a/backend/internal/application/conversation/service_media_video.go +++ b/backend/internal/application/conversation/service_media_video.go @@ -96,6 +96,7 @@ func (s *Service) StreamMediaVideo(ctx context.Context, input MediaVideoInput) ( if !llm.IsVideoGenerationAdapter(route.Protocol) { return nil, ErrMediaRouteProtocolMismatch } + videoEndpoint := llm.DefaultEndpointForAdapter(route.Protocol) if strings.TrimSpace(conversation.Model) != strings.TrimSpace(route.PlatformModelName) { conversation.Model = strings.TrimSpace(route.PlatformModelName) conversation.Provider = inferProvider(conversation.Model) @@ -115,7 +116,7 @@ func (s *Service) StreamMediaVideo(ctx context.Context, input MediaVideoInput) ( UserID: input.UserID, ConversationID: input.ConversationID, TaskType: channel.TaskTypeVideoGeneration, - Endpoint: llm.EndpointInteractions, + Endpoint: videoEndpoint, Provider: strings.TrimSpace(conversation.Provider), ProviderProtocol: route.Protocol, UpstreamID: route.UpstreamID, @@ -227,7 +228,7 @@ func (s *Service) StreamMediaVideo(ctx context.Context, input MediaVideoInput) ( ConnectTimeoutMS: route.ConnectTimeoutMS, ReadTimeoutMS: route.ReadTimeoutMS, StreamIdleTimeoutMS: route.StreamIdleTimeoutMS, - Endpoint: llm.EndpointInteractions, + Endpoint: videoEndpoint, UpstreamModel: route.UpstreamModel, AttributionReferer: attributionReferer, AttributionTitle: attributionTitle, diff --git a/backend/internal/application/settings/model_option_policy.go b/backend/internal/application/settings/model_option_policy.go index 606c50067..5c9bdb16e 100644 --- a/backend/internal/application/settings/model_option_policy.go +++ b/backend/internal/application/settings/model_option_policy.go @@ -20,6 +20,7 @@ var validModelOptionProtocolKeys = map[string]struct{}{ "xai_responses": {}, "xai_image": {}, "xai_image_edits": {}, + "xai_video": {}, "gemini_generate_content": {}, "google_image_generation": {}, "gemini_interactions": {}, diff --git a/backend/internal/application/settings/service.go b/backend/internal/application/settings/service.go index 4fbb6b07e..87574d209 100644 --- a/backend/internal/application/settings/service.go +++ b/backend/internal/application/settings/service.go @@ -198,12 +198,16 @@ func isLegacyDefaultModelOptionAllowedPaths(value string) bool { if err := json.Unmarshal([]byte(strings.TrimSpace(value)), ¤t); err != nil { return false } - legacy := map[string][]string{} - if err := json.Unmarshal([]byte(config.DefaultModelOptionAllowedPathsJSON()), &legacy); err != nil { + previousDefault := map[string][]string{} + if err := json.Unmarshal([]byte(config.DefaultModelOptionAllowedPathsJSON()), &previousDefault); err != nil { return false } - legacy["xai_responses"] = []string{"reasoning.effort"} - return sameStringSliceMap(current, legacy) + delete(previousDefault, "xai_video") + if sameStringSliceMap(current, previousDefault) { + return true + } + previousDefault["xai_responses"] = []string{"reasoning.effort"} + return sameStringSliceMap(current, previousDefault) } func sameStringSliceMap(left map[string][]string, right map[string][]string) bool { diff --git a/backend/internal/application/settings/service_seed_test.go b/backend/internal/application/settings/service_seed_test.go index dc57ad65d..22662adbd 100644 --- a/backend/internal/application/settings/service_seed_test.go +++ b/backend/internal/application/settings/service_seed_test.go @@ -135,6 +135,7 @@ func TestSeedMigratesLegacyDefaultModelOptionAllowedPaths(t *testing.T) { if err := json.Unmarshal([]byte(config.DefaultModelOptionAllowedPathsJSON()), &legacy); err != nil { t.Fatalf("decode current model option defaults: %v", err) } + delete(legacy, "xai_video") legacy["xai_responses"] = []string{"reasoning.effort"} legacyJSON, err := json.Marshal(legacy) if err != nil { @@ -157,6 +158,32 @@ func TestSeedMigratesLegacyDefaultModelOptionAllowedPaths(t *testing.T) { } } +func TestSeedAddsXAIVideoToPreviousDefaultModelOptionAllowedPaths(t *testing.T) { + previousDefault := map[string][]string{} + if err := json.Unmarshal([]byte(config.DefaultModelOptionAllowedPathsJSON()), &previousDefault); err != nil { + t.Fatalf("decode current model option defaults: %v", err) + } + delete(previousDefault, "xai_video") + previousJSON, err := json.Marshal(previousDefault) + if err != nil { + t.Fatalf("encode previous model option defaults: %v", err) + } + repo := newSettingsSeedRepo(domainsettings.SystemSetting{ + Namespace: "chat", + Key: "model_option_allowed_paths", + Value: string(previousJSON), + ValueType: "json", + }) + service := NewService(repo, "") + + if err := service.Seed(context.Background(), config.Config{}); err != nil { + t.Fatalf("seed settings: %v", err) + } + if got := repo.items["chat:model_option_allowed_paths"].Value; got != config.DefaultModelOptionAllowedPathsJSON() { + t.Fatalf("expected xAI video defaults to be added, got %q", got) + } +} + func TestSeedKeepsCustomModelOptionAllowedPaths(t *testing.T) { custom := `{"default":["temperature"],"xai_responses":["reasoning.effort"]}` repo := newSettingsSeedRepo(domainsettings.SystemSetting{ diff --git a/backend/internal/infra/config/config.go b/backend/internal/infra/config/config.go index f753c4428..209f7a570 100644 --- a/backend/internal/infra/config/config.go +++ b/backend/internal/infra/config/config.go @@ -159,6 +159,11 @@ func DefaultModelOptionAllowedPathsJSON() string { "resolution", "response_format" ], + "xai_video": [ + "aspect_ratio", + "duration", + "resolution" + ], "gemini_generate_content": [ "generationConfig.temperature", "generationConfig.topP", diff --git a/backend/internal/infra/llm/adapter.go b/backend/internal/infra/llm/adapter.go index 69fee82af..62e23b01d 100644 --- a/backend/internal/infra/llm/adapter.go +++ b/backend/internal/infra/llm/adapter.go @@ -22,6 +22,7 @@ const ( AdapterXAIResponses = "xai_responses" // POST /v1/responses(OpenAI 兼容) AdapterXAIImage = "xai_image" // POST /v1/images/generations AdapterXAIImageEdits = "xai_image_edits" // POST /v1/images/edits + AdapterXAIVideo = "xai_video" // POST /v1/videos/generations + GET /v1/videos/{request_id} ) var ( @@ -62,7 +63,8 @@ func IsKnownAdapter(raw string) bool { AdapterGeminiInteractions, AdapterXAIResponses, AdapterXAIImage, - AdapterXAIImageEdits: + AdapterXAIImageEdits, + AdapterXAIVideo: return true default: return false @@ -73,7 +75,7 @@ func IsKnownAdapter(raw string) bool { func IsImplementedAdapter(raw string) bool { switch NormalizeAdapter(raw) { case AdapterOpenAIResponses, AdapterOpenRouterChat, AdapterOpenRouterResponses, AdapterOpenAIChatCompletions, AdapterOpenAIImageGenerations, AdapterOpenAIImageEdits, AdapterXAIResponses, - AdapterAnthropicMessages, AdapterGoogleGenerateContent, AdapterGoogleImageGeneration, AdapterGeminiInteractions, AdapterXAIImage, AdapterXAIImageEdits: + AdapterAnthropicMessages, AdapterGoogleGenerateContent, AdapterGoogleImageGeneration, AdapterGeminiInteractions, AdapterXAIImage, AdapterXAIImageEdits, AdapterXAIVideo: return true default: return false @@ -138,7 +140,12 @@ func IsImageEditAdapter(raw string) bool { // IsVideoGenerationAdapter 返回协议是否属于独立视频生成链路。 func IsVideoGenerationAdapter(raw string) bool { - return NormalizeAdapter(raw) == AdapterGeminiInteractions + switch NormalizeAdapter(raw) { + case AdapterGeminiInteractions, AdapterXAIVideo: + return true + default: + return false + } } // DefaultEndpointForAdapter 返回协议对应的固定端点标识。 @@ -150,6 +157,8 @@ func DefaultEndpointForAdapter(adapter string) string { return EndpointImageGenerations case AdapterOpenAIImageEdits, AdapterXAIImageEdits: return EndpointImageEdits + case AdapterXAIVideo: + return EndpointVideoGenerations case AdapterGeminiInteractions: return EndpointInteractions default: diff --git a/backend/internal/infra/llm/adapter_test.go b/backend/internal/infra/llm/adapter_test.go index 2313d9d94..a5f41da96 100644 --- a/backend/internal/infra/llm/adapter_test.go +++ b/backend/internal/infra/llm/adapter_test.go @@ -69,3 +69,18 @@ func TestImageAdapterCapabilities(t *testing.T) { t.Fatalf("expected xAI image edits protocol to support image editing") } } + +func TestXAIVideoAdapterCapabilities(t *testing.T) { + if !IsKnownAdapter(AdapterXAIVideo) || !IsImplementedAdapter(AdapterXAIVideo) { + t.Fatalf("expected xAI video adapter to be known and implemented") + } + if !IsVideoGenerationAdapter(AdapterXAIVideo) { + t.Fatalf("expected xAI video adapter to support video generation") + } + if SupportsStreamingAdapter(AdapterXAIVideo) { + t.Fatalf("expected xAI video adapter to use asynchronous polling instead of streaming") + } + if got := DefaultEndpointForAdapter(AdapterXAIVideo); got != EndpointVideoGenerations { + t.Fatalf("expected xAI video endpoint, got %q", got) + } +} diff --git a/backend/internal/infra/llm/client.go b/backend/internal/infra/llm/client.go index d45517bc9..b121dc8a6 100644 --- a/backend/internal/infra/llm/client.go +++ b/backend/internal/infra/llm/client.go @@ -30,6 +30,8 @@ const ( EndpointImageGenerations = "image_generations" // EndpointImageEdits 表示 OpenAI Images API 编辑端点。 EndpointImageEdits = "image_edits" + // EndpointVideoGenerations 表示异步视频生成端点。 + EndpointVideoGenerations = "video_generations" // EndpointInteractions 表示 Gemini Interactions API 端点。 EndpointInteractions = "interactions" ) @@ -813,6 +815,7 @@ func NewClient(outboundPolicy security.OutboundPolicy) *Client { AdapterXAIResponses: &xAIResponsesAdapter{client: client}, AdapterXAIImage: &xAIImageAdapter{client: client}, AdapterXAIImageEdits: &xAIImageEditsAdapter{client: client}, + AdapterXAIVideo: &xAIVideoAdapter{client: client}, AdapterAnthropicMessages: &anthropicMessagesAdapter{client: client}, AdapterGoogleGenerateContent: &geminiGenerateContentAdapter{client: client}, AdapterGoogleImageGeneration: &geminiImageGenerationAdapter{client: client}, @@ -1637,6 +1640,8 @@ func normalizeEndpoint(raw string) string { return EndpointImageGenerations case EndpointImageEdits: return EndpointImageEdits + case EndpointVideoGenerations: + return EndpointVideoGenerations case EndpointInteractions: return EndpointInteractions default: diff --git a/backend/internal/infra/llm/endpoint_url_test.go b/backend/internal/infra/llm/endpoint_url_test.go index 33e9f48ec..dd26e3cd1 100644 --- a/backend/internal/infra/llm/endpoint_url_test.go +++ b/backend/internal/infra/llm/endpoint_url_test.go @@ -45,6 +45,12 @@ func TestBuildOpenAICompatibleURLsRespectVersionedBasePath(t *testing.T) { endpoint: EndpointImageGenerations, want: "https://api.x.ai/v1/images/generations", }, + { + name: "xai video generations endpoint", + baseURL: "https://api.x.ai/v1", + endpoint: EndpointVideoGenerations, + want: "https://api.x.ai/v1/videos/generations", + }, { name: "xai proxy plain base gets v1 image endpoint", baseURL: "https://proxy.example.com", diff --git a/backend/internal/infra/llm/openai.go b/backend/internal/infra/llm/openai.go index a8d1d0ea2..b8abec48f 100644 --- a/backend/internal/infra/llm/openai.go +++ b/backend/internal/infra/llm/openai.go @@ -364,6 +364,8 @@ func buildOpenAIRequestURL(baseURL string, endpoint string) string { return buildVersionedEndpointURL(baseURL, "v1", "/images/generations") case EndpointImageEdits: return buildVersionedEndpointURL(baseURL, "v1", "/images/edits") + case EndpointVideoGenerations: + return buildVersionedEndpointURL(baseURL, "v1", "/videos/generations") default: return buildVersionedEndpointURL(baseURL, "v1", "/responses") } diff --git a/backend/internal/infra/llm/xai_images.go b/backend/internal/infra/llm/xai_images.go index 75fd7a5ba..320a8df72 100644 --- a/backend/internal/infra/llm/xai_images.go +++ b/backend/internal/infra/llm/xai_images.go @@ -6,6 +6,7 @@ import ( "encoding/base64" "encoding/json" "fmt" + "math" "net/http" "strings" ) @@ -109,7 +110,7 @@ func (c *Client) generateXAIImage(ctx context.Context, route RouteConfig, input } setAdditionalHeaders(req, route.HeadersJSON) - resp, err := c.doRouteRequest(route, req) + resp, err := c.doRouteGenerationRequest(route, req) if err != nil { return nil, err } @@ -117,13 +118,22 @@ func (c *Client) generateXAIImage(ctx context.Context, route RouteConfig, input body, err := readUpstreamBody(resp.Body) if err != nil { + if resp.StatusCode >= 200 && resp.StatusCode < 300 { + return nil, MarkRequestAccepted(attachUpstreamDebug(err, upstreamDebugSnapshot(req, debugPayload, resp, body))) + } return nil, err } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, parseUpstreamError(resp.StatusCode, body, upstreamDebugSnapshot(req, debugPayload, resp, body)) } - return parseXAIImageOutput(body, modelParamString(input.Options, "response_format"), protocol) + debug := upstreamDebugSnapshot(req, debugPayload, resp, body) + output, err := parseXAIImageOutput(body, protocol) + if err != nil { + return nil, MarkRequestAccepted(attachUpstreamDebug(err, debug)) + } + output.Debug = debug + return output, nil } // buildXAIImageRequest 根据任务端点构造 xAI 图片生成或编辑请求。 @@ -173,7 +183,7 @@ func buildXAIImageEditRequestBody(model string, input GenerateInput) (map[string if len(imageInputs) == 1 { payload["image"] = imageInputs[0] } else { - payload["image"] = imageInputs + payload["images"] = imageInputs } applyXAIImageParams(payload, input.Options) debugBody, _ := json.Marshal(map[string]interface{}{ @@ -198,19 +208,67 @@ func xAIImageURLPayload(image ContentPart) map[string]interface{} { // applyXAIImageParams 从 options 中提取 xAI 图片生成官方参数。 func applyXAIImageParams(payload map[string]interface{}, options map[string]interface{}) { - payload["response_format"] = defaultImageResponseFormat(options) - for _, key := range []string{"aspect_ratio", "resolution"} { - if value := modelParamString(options, key); value != "" { - payload[key] = value - } + if value := xAIImageResponseFormat(options); value != "" { + payload["response_format"] = value } - if value := modelParamInt(options, "n"); value > 0 { + if value := strings.ToLower(modelParamString(options, "aspect_ratio")); isXAIImageAspectRatio(value) { + payload["aspect_ratio"] = value + } + if value := strings.ToLower(modelParamString(options, "resolution")); isXAIImageResolution(value) { + payload["resolution"] = value + } + if value, ok := xAIMediaIntegerOption(options, "n"); ok && value >= 1 && value <= 10 { payload["n"] = value } } +func xAIMediaIntegerOption(options map[string]interface{}, key string) (int, bool) { + if options == nil { + return 0, false + } + value, ok := options[key] + if !ok { + return 0, false + } + switch typed := value.(type) { + case int: + return typed, true + case int64: + return int(typed), true + case float64: + if math.Trunc(typed) == typed { + return int(typed), true + } + } + return 0, false +} + +func isXAIImageAspectRatio(value string) bool { + switch value { + case "1:1", "3:4", "4:3", "9:16", "16:9", "2:3", "3:2", "9:19.5", "19.5:9", "9:20", "20:9", "1:2", "2:1", "auto": + return true + default: + return false + } +} + +func isXAIImageResolution(value string) bool { + return value == "1k" || value == "2k" +} + +func xAIImageResponseFormat(options map[string]interface{}) string { + switch strings.ToLower(modelParamString(options, "response_format")) { + case "url": + return "url" + case "b64_json": + return "b64_json" + default: + return "" + } +} + // parseXAIImageOutput 解析 xAI 图片响应;图片字节只进入 GeneratedImages。 -func parseXAIImageOutput(body []byte, responseFormat string, protocol string) (*GenerateOutput, error) { +func parseXAIImageOutput(body []byte, protocol string) (*GenerateOutput, error) { parsed := make(map[string]interface{}) if err := json.Unmarshal(body, &parsed); err != nil { return nil, err @@ -227,7 +285,7 @@ func parseXAIImageOutput(body []byte, responseFormat string, protocol string) (* data := asSlice(parsed["data"]) citations := make([]string, 0, len(data)) for _, item := range data { - if image, ok := parseXAIImagePayload(asMap(item), responseFormat); ok { + if image, ok := parseXAIImagePayload(asMap(item)); ok { if url := strings.TrimSpace(image.URL); url != "" { citations = append(citations, url) } @@ -235,7 +293,7 @@ func parseXAIImageOutput(body []byte, responseFormat string, protocol string) (* } } if len(data) == 0 { - if image, ok := parseXAIImagePayload(parsed, responseFormat); ok { + if image, ok := parseXAIImagePayload(parsed); ok { if url := strings.TrimSpace(image.URL); url != "" { citations = append(citations, url) } @@ -246,10 +304,11 @@ func parseXAIImageOutput(body []byte, responseFormat string, protocol string) (* return result, nil } -func parseXAIImagePayload(payload map[string]interface{}, responseFormat string) (GeneratedImage, bool) { +func parseXAIImagePayload(payload map[string]interface{}) (GeneratedImage, bool) { if len(payload) == 0 { return GeneratedImage{}, false } + mimeType := xAIImageMIMEType(payload) revisedPrompt := strings.TrimSpace(getString(payload["revised_prompt"])) if revisedPrompt == "" { revisedPrompt = strings.TrimSpace(getString(payload["revisedPrompt"])) @@ -257,26 +316,31 @@ func parseXAIImagePayload(payload map[string]interface{}, responseFormat string) if url := strings.TrimSpace(getString(payload["url"])); url != "" { return GeneratedImage{ URL: url, - MIMEType: xAIImageMIMEType(responseFormat), + MIMEType: mimeType, RevisedPrompt: revisedPrompt, }, true } if b64 := strings.TrimSpace(getString(payload["b64_json"])); b64 != "" { return GeneratedImage{ B64JSON: b64, - MIMEType: xAIImageMIMEType(responseFormat), + MIMEType: mimeType, + RevisedPrompt: revisedPrompt, + }, true + } + if publicURL := strings.TrimSpace(getString(asMap(payload["file_output"])["public_url"])); publicURL != "" { + return GeneratedImage{ + URL: publicURL, + MIMEType: mimeType, RevisedPrompt: revisedPrompt, }, true } return GeneratedImage{}, false } -// xAIImageMIMEType 根据 xAI 文档示例的默认图片格式给 base64 结果设置初始 MIME。 -func xAIImageMIMEType(responseFormat string) string { - switch strings.TrimSpace(strings.ToLower(responseFormat)) { - case "b64_json", "url", "": - return "image/jpeg" - default: - return "image/jpeg" +// xAIImageMIMEType 优先采用官方响应中的 MIME;旧代理未返回时回退为 JPEG。 +func xAIImageMIMEType(payload map[string]interface{}) string { + if mimeType := strings.ToLower(strings.TrimSpace(getString(payload["mime_type"]))); strings.HasPrefix(mimeType, "image/") { + return mimeType } + return "image/jpeg" } diff --git a/backend/internal/infra/llm/xai_images_test.go b/backend/internal/infra/llm/xai_images_test.go index 8b7b58b97..5d3abd27b 100644 --- a/backend/internal/infra/llm/xai_images_test.go +++ b/backend/internal/infra/llm/xai_images_test.go @@ -41,15 +41,34 @@ func TestBuildXAIImageRequestBody(t *testing.T) { } } -func TestBuildXAIImageRequestBodyDefaultsToBase64(t *testing.T) { +func TestBuildXAIImageRequestBodyDropsUnsupportedParams(t *testing.T) { + payload, err := buildXAIImageRequestBody("grok-imagine-image-quality", GenerateInput{ + Messages: []Message{{Role: "user", Content: "A clean product render"}}, + Options: map[string]interface{}{ + "aspect_ratio": "21:9", + "n": 2.5, + "resolution": "4k", + }, + }) + if err != nil { + t.Fatalf("build xAI image request body: %v", err) + } + for _, key := range []string{"aspect_ratio", "n", "resolution"} { + if _, ok := payload[key]; ok { + t.Fatalf("unsupported xAI image param %q must be removed: %#v", key, payload) + } + } +} + +func TestBuildXAIImageRequestBodyPreservesOfficialDefaultResponseFormat(t *testing.T) { payload, err := buildXAIImageRequestBody("grok-imagine-image-quality", GenerateInput{ Messages: []Message{{Role: "user", Content: "A clean product render"}}, }) if err != nil { t.Fatalf("build xAI image request body: %v", err) } - if payload["response_format"] != "b64_json" { - t.Fatalf("expected xAI image generation to default to base64, got %#v", payload) + if _, ok := payload["response_format"]; ok { + t.Fatalf("expected xAI to apply its documented URL default, got %#v", payload) } } @@ -78,8 +97,8 @@ func TestBuildXAIImageEditRequestBody(t *testing.T) { if payload["aspect_ratio"] != "1:1" || payload["resolution"] != "2k" { t.Fatalf("expected xAI edit params, got %#v", payload) } - if payload["response_format"] != "b64_json" { - t.Fatalf("expected xAI image edit to default to base64, got %#v", payload) + if _, ok := payload["response_format"]; ok { + t.Fatalf("expected xAI to apply its documented URL default, got %#v", payload) } image := payload["image"].(map[string]interface{}) if image["type"] != "image_url" { @@ -110,7 +129,10 @@ func TestBuildXAIImageEditRequestBodyAllowsUpToThreeImages(t *testing.T) { if err != nil { t.Fatalf("build xAI multi-image edit request body: %v", err) } - images := payload["image"].([]map[string]interface{}) + if _, ok := payload["image"]; ok { + t.Fatalf("multi-reference edit must not send the singular image field: %#v", payload) + } + images := payload["images"].([]map[string]interface{}) if len(images) != 3 { t.Fatalf("expected three ordered image inputs, got %#v", images) } @@ -252,10 +274,10 @@ func TestParseXAIImageOutput(t *testing.T) { output, err := parseXAIImageOutput([]byte(`{ "id": "img_xai_1", "data": [ - {"url": "https://example.com/a.jpg"}, - {"b64_json": "aGVsbG8=", "revised_prompt": "A revised render"} + {"url": "https://example.com/a.jpg", "mime_type": "image/webp"}, + {"b64_json": "aGVsbG8=", "mime_type": "image/png", "revised_prompt": "A revised render"} ] - }`), "b64_json", AdapterXAIImage) + }`), AdapterXAIImage) if err != nil { t.Fatalf("parse xAI image output: %v", err) } @@ -265,10 +287,10 @@ func TestParseXAIImageOutput(t *testing.T) { if len(output.GeneratedImages) != 2 { t.Fatalf("expected two generated images, got %#v", output.GeneratedImages) } - if output.GeneratedImages[0].URL != "https://example.com/a.jpg" { + if output.GeneratedImages[0].URL != "https://example.com/a.jpg" || output.GeneratedImages[0].MIMEType != "image/webp" { t.Fatalf("unexpected URL image: %#v", output.GeneratedImages[0]) } - if output.GeneratedImages[1].B64JSON != "aGVsbG8=" || output.GeneratedImages[1].MIMEType != "image/jpeg" { + if output.GeneratedImages[1].B64JSON != "aGVsbG8=" || output.GeneratedImages[1].MIMEType != "image/png" { t.Fatalf("unexpected base64 image: %#v", output.GeneratedImages[1]) } if len(output.Citations) != 1 || output.Citations[0] != "https://example.com/a.jpg" { diff --git a/backend/internal/infra/llm/xai_videos.go b/backend/internal/infra/llm/xai_videos.go new file mode 100644 index 000000000..781ffd330 --- /dev/null +++ b/backend/internal/infra/llm/xai_videos.go @@ -0,0 +1,361 @@ +package llm + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "fmt" + "maps" + "net/http" + "net/url" + "strconv" + "strings" + "time" +) + +const defaultXAIVideoPollInterval = time.Second + +// xAIVideoAdapter 实现 xAI 异步视频生成协议。 +type xAIVideoAdapter struct { + client *Client +} + +func (a *xAIVideoAdapter) Name() string { return AdapterXAIVideo } + +func (a *xAIVideoAdapter) Generate(ctx context.Context, route RouteConfig, input GenerateInput) (*GenerateOutput, error) { + route.Protocol = AdapterXAIVideo + route.Endpoint = EndpointVideoGenerations + return a.client.generateXAIVideo(ctx, route, input) +} + +func (a *xAIVideoAdapter) GenerateStream( + context.Context, + RouteConfig, + GenerateInput, + func(GenerateStreamEvent) error, +) (*GenerateOutput, error) { + return nil, fmt.Errorf("%w: %s", ErrUnsupportedStream, AdapterXAIVideo) +} + +func (a *xAIVideoAdapter) ListModels(ctx context.Context, route RouteConfig) ([]ModelItem, error) { + route.Protocol = AdapterXAIVideo + return a.client.listModelsOpenAICompatible(ctx, route) +} + +// generateXAIVideo 提交视频任务,并在同一请求超时范围内轮询官方结果端点。 +func (c *Client) generateXAIVideo(ctx context.Context, route RouteConfig, input GenerateInput) (*GenerateOutput, error) { + requestBody, debugBody, err := buildXAIVideoRequestBody(route.UpstreamModel, input) + if err != nil { + return nil, err + } + payload, err := json.Marshal(requestBody) + if err != nil { + return nil, err + } + requestURL := buildOpenAIRequestURL(route.BaseURL, EndpointVideoGenerations) + if requestURL == "" { + return nil, fmt.Errorf("invalid base url") + } + + requestCtx, cancel := context.WithTimeout(ctx, resolveReadTimeout(route.ReadTimeoutMS)) + defer cancel() + + req, err := newXAIMediaRequest(requestCtx, http.MethodPost, requestURL, payload, route) + if err != nil { + return nil, err + } + resp, err := c.doRouteGenerationRequest(route, req) + if err != nil { + return nil, err + } + body, readErr := readUpstreamBody(resp.Body) + _ = resp.Body.Close() + if readErr != nil { + if resp.StatusCode >= 200 && resp.StatusCode < 300 { + debug := upstreamDebugSnapshot(req, debugBody, resp, body) + return nil, acceptedXAIVideoResponseError(readErr, debug) + } + return nil, readErr + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return nil, parseUpstreamError(resp.StatusCode, body, upstreamDebugSnapshot(req, debugBody, resp, body)) + } + + requestID, err := parseXAIVideoRequestID(body) + if err != nil { + return nil, acceptedXAIVideoResponseError(err, upstreamDebugSnapshot(req, debugBody, resp, body)) + } + return c.pollXAIVideoResult(requestCtx, route, requestID) +} + +func newXAIMediaRequest(ctx context.Context, method string, requestURL string, payload []byte, route RouteConfig) (*http.Request, error) { + req, err := http.NewRequestWithContext(ctx, method, requestURL, bytes.NewReader(payload)) + if err != nil { + return nil, err + } + if len(payload) > 0 { + req.Header.Set("Content-Type", "application/json") + } + if apiKey := strings.TrimSpace(route.APIKey); apiKey != "" { + req.Header.Set("Authorization", "Bearer "+apiKey) + } + setAdditionalHeaders(req, route.HeadersJSON) + return req, nil +} + +func buildXAIVideoRequestBody(model string, input GenerateInput) (map[string]interface{}, []byte, error) { + prompt := buildOpenAIImageGenerationPrompt(input.Messages) + images := collectImageInputParts(input.Messages) + if strings.TrimSpace(prompt) == "" && len(images) == 0 { + return nil, nil, fmt.Errorf("video generation prompt or input image required") + } + if len(images) > 1 { + return nil, nil, fmt.Errorf("too many video generation input images") + } + + payload := map[string]interface{}{ + "model": strings.TrimSpace(model), + } + if strings.TrimSpace(prompt) != "" { + payload["prompt"] = strings.TrimSpace(prompt) + } + if len(images) == 1 { + payload["image"] = xAIVideoImagePayload(images[0]) + } + applyXAIVideoParams(payload, input.Options) + + debugPayload := make(map[string]interface{}, len(payload)+1) + for key, value := range payload { + if key != "image" { + debugPayload[key] = value + } + } + debugPayload["image_count"] = len(images) + debugBody, _ := json.Marshal(debugPayload) + return payload, debugBody, nil +} + +func xAIVideoImagePayload(image ContentPart) map[string]interface{} { + mimeType := strings.ToLower(strings.TrimSpace(image.MimeType)) + switch mimeType { + case "image/png", "image/webp", "image/jpeg": + default: + mimeType = "image/jpeg" + } + return map[string]interface{}{ + "url": "data:" + mimeType + ";base64," + base64.StdEncoding.EncodeToString(image.Data), + } +} + +func applyXAIVideoParams(payload map[string]interface{}, options map[string]interface{}) { + normalized := maps.Clone(options) + SanitizeXAIVideoOptions(normalized) + for _, key := range []string{"aspect_ratio", "duration", "resolution"} { + if value, ok := normalized[key]; ok { + payload[key] = value + } + } +} + +// SanitizeXAIVideoOptions 将 xAI 视频协议参数收敛为实际会上送的规范值。 +// Application 层复用该函数,保证有效参数、计费和 adapter 请求一致。 +func SanitizeXAIVideoOptions(options map[string]interface{}) { + if len(options) == 0 { + return + } + aspectRatio := strings.ToLower(modelParamString(options, "aspect_ratio")) + if isXAIVideoAspectRatio(aspectRatio) { + options["aspect_ratio"] = aspectRatio + } else { + delete(options, "aspect_ratio") + } + duration, durationOK := xAIMediaIntegerOption(options, "duration") + if durationOK && duration >= 1 && duration <= 15 { + options["duration"] = duration + } else { + delete(options, "duration") + } + resolution := strings.ToLower(modelParamString(options, "resolution")) + if isXAIVideoResolution(resolution) { + options["resolution"] = resolution + } else { + delete(options, "resolution") + } +} + +func isXAIVideoAspectRatio(value string) bool { + switch value { + case "1:1", "16:9", "9:16", "4:3", "3:4", "3:2", "2:3": + return true + default: + return false + } +} + +func isXAIVideoResolution(value string) bool { + switch value { + case "480p", "720p", "1080p": + return true + default: + return false + } +} + +func parseXAIVideoRequestID(body []byte) (string, error) { + parsed := make(map[string]interface{}) + if err := json.Unmarshal(body, &parsed); err != nil { + return "", err + } + requestID := strings.TrimSpace(getString(parsed["request_id"])) + if requestID == "" { + return "", fmt.Errorf("xAI video response missing request_id") + } + return requestID, nil +} + +func (c *Client) pollXAIVideoResult(ctx context.Context, route RouteConfig, requestID string) (*GenerateOutput, error) { + requestURL := buildXAIVideoResultURL(route.BaseURL, requestID) + if requestURL == "" { + return nil, MarkRequestAccepted(fmt.Errorf("invalid xAI video result url")) + } + + for { + req, err := newXAIMediaRequest(ctx, http.MethodGet, requestURL, nil, route) + if err != nil { + return nil, MarkRequestAccepted(err) + } + resp, err := c.doRouteRequest(route, req) + if err != nil { + return nil, MarkRequestAccepted(err) + } + body, readErr := readUpstreamBody(resp.Body) + _ = resp.Body.Close() + debug := upstreamDebugSnapshot(req, nil, resp, body) + if readErr != nil { + return nil, acceptedXAIVideoResponseError(readErr, debug) + } + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusAccepted { + return nil, MarkRequestAccepted(parseUpstreamError(resp.StatusCode, body, debug)) + } + + output, pending, err := parseXAIVideoResult(body, requestID) + if err != nil { + return nil, acceptedXAIVideoResponseError(err, debug) + } + if !pending { + output.Debug = debug + return output, nil + } + if err := waitXAIVideoPoll(ctx, xAIVideoPollDelay(resp.Header.Get("Retry-After"))); err != nil { + return nil, MarkRequestAccepted(err) + } + } +} + +func buildXAIVideoResultURL(baseURL string, requestID string) string { + id := strings.TrimSpace(requestID) + if id == "" { + return "" + } + return buildVersionedEndpointURL(baseURL, "v1", "/videos/"+url.PathEscape(id)) +} + +func parseXAIVideoResult(body []byte, requestID string) (*GenerateOutput, bool, error) { + parsed := make(map[string]interface{}) + if err := json.Unmarshal(body, &parsed); err != nil { + return nil, false, err + } + status := strings.ToLower(strings.TrimSpace(getString(parsed["status"]))) + switch status { + case "pending": + return nil, true, nil + case "failed": + errorPayload := asMap(parsed["error"]) + code := strings.TrimSpace(getString(errorPayload["code"])) + message := strings.TrimSpace(getString(errorPayload["message"])) + if message == "" { + message = "xAI video generation failed" + } + if code != "" { + message = code + ": " + message + } + return nil, false, fmt.Errorf("xAI video generation failed: %s", message) + case "done": + default: + return nil, false, fmt.Errorf("unexpected xAI video status %q", status) + } + + videoPayload := asMap(parsed["video"]) + if respectsModeration, ok := videoPayload["respect_moderation"].(bool); ok && !respectsModeration { + return nil, false, fmt.Errorf("xAI video result was blocked by content moderation") + } + videoURL := strings.TrimSpace(getString(videoPayload["url"])) + fileOutput := asMap(videoPayload["file_output"]) + if videoURL == "" { + videoURL = strings.TrimSpace(getString(fileOutput["public_url"])) + } + if videoURL == "" { + return nil, false, fmt.Errorf("xAI video result missing downloadable URL") + } + result := &GenerateOutput{ + ResponseID: strings.TrimSpace(requestID), + ToolCalls: make([]ToolCall, 0), + ServerToolCalls: make([]ToolCall, 0), + GeneratedVideos: []GeneratedVideo{{ + URL: videoURL, + MIMEType: "video/mp4", + FileName: strings.TrimSpace(getString(fileOutput["filename"])), + }}, + RawJSON: string(body), + } + result.Usage.RawUsageJSON = rawUsageJSONFromPath(parsed, "usage") + return result, false, nil +} + +func acceptedXAIVideoResponseError(err error, debug *UpstreamDebugSnapshot) error { + if err == nil { + return nil + } + body := "" + if debug != nil { + body = debug.Response.Body + } + return MarkRequestAccepted(&UpstreamError{ + StatusCode: http.StatusBadGateway, + Message: strings.TrimSpace(err.Error()), + Body: body, + Debug: debug, + }) +} + +func xAIVideoPollDelay(retryAfter string) time.Duration { + seconds, err := strconv.Atoi(strings.TrimSpace(retryAfter)) + if err != nil || seconds < 0 { + return defaultXAIVideoPollInterval + } + delay := time.Duration(seconds) * time.Second + if delay > 10*time.Second { + return 10 * time.Second + } + return delay +} + +func waitXAIVideoPoll(ctx context.Context, delay time.Duration) error { + if delay <= 0 { + select { + case <-ctx.Done(): + return ctx.Err() + default: + return nil + } + } + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} diff --git a/backend/internal/infra/llm/xai_videos_test.go b/backend/internal/infra/llm/xai_videos_test.go new file mode 100644 index 000000000..b912c2641 --- /dev/null +++ b/backend/internal/infra/llm/xai_videos_test.go @@ -0,0 +1,197 @@ +package llm + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestBuildXAIVideoRequestBody(t *testing.T) { + payload, debugBody, err := buildXAIVideoRequestBody("grok-imagine-video", GenerateInput{ + Messages: []Message{{ + Role: "user", + Parts: []ContentPart{ + {Kind: ContentPartText, Text: "Animate the scene"}, + {Kind: ContentPartImage, MimeType: "image/png", Data: []byte("source")}, + }, + }}, + Options: map[string]interface{}{ + "aspect_ratio": "16:9", + "duration": 6, + "resolution": "720P", + "prompt": "must not override messages", + "output": "must not pass through", + }, + }) + if err != nil { + t.Fatalf("build xAI video request body: %v", err) + } + if payload["model"] != "grok-imagine-video" || payload["prompt"] != "Animate the scene" { + t.Fatalf("unexpected model or prompt: %#v", payload) + } + if payload["aspect_ratio"] != "16:9" || payload["duration"] != 6 || payload["resolution"] != "720p" { + t.Fatalf("expected documented xAI video params, got %#v", payload) + } + image := asMap(payload["image"]) + if !strings.HasPrefix(getString(image["url"]), "data:image/png;base64,c291cmNl") { + t.Fatalf("expected data URL image input, got %#v", image) + } + for _, key := range []string{"output"} { + if _, ok := payload[key]; ok { + t.Fatalf("unexpected xAI video param %q in %#v", key, payload) + } + } + if strings.Contains(string(debugBody), "c291cmNl") || !strings.Contains(string(debugBody), `"image_count":1`) { + t.Fatalf("debug body must summarize without source bytes: %s", debugBody) + } +} + +func TestBuildXAIVideoRequestBodyRejectsMultipleImages(t *testing.T) { + _, _, err := buildXAIVideoRequestBody("grok-imagine-video", GenerateInput{ + Messages: []Message{{ + Role: "user", + Parts: []ContentPart{ + {Kind: ContentPartText, Text: "Animate the scene"}, + {Kind: ContentPartImage, MimeType: "image/png", Data: []byte("one")}, + {Kind: ContentPartImage, MimeType: "image/png", Data: []byte("two")}, + }, + }}, + }) + if err == nil || !strings.Contains(err.Error(), "too many") { + t.Fatalf("expected multiple image validation error, got %v", err) + } +} + +func TestBuildXAIVideoRequestBodyDropsUnsupportedParams(t *testing.T) { + payload, _, err := buildXAIVideoRequestBody("grok-imagine-video", GenerateInput{ + Messages: []Message{{Role: "user", Content: "Animate the scene"}}, + Options: map[string]interface{}{ + "aspect_ratio": "21:9", + "duration": 8.5, + "resolution": "4k", + }, + }) + if err != nil { + t.Fatalf("build xAI video request body: %v", err) + } + for _, key := range []string{"aspect_ratio", "duration", "resolution"} { + if _, ok := payload[key]; ok { + t.Fatalf("unsupported xAI video param %q must be removed: %#v", key, payload) + } + } +} + +func TestGenerateXAIVideoSubmitsAndPolls(t *testing.T) { + postCount := 0 + pollCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got := r.Header.Get("Authorization"); got != "Bearer xai-key" { + t.Fatalf("unexpected auth header %q", got) + } + switch { + case r.Method == http.MethodPost && r.URL.Path == "/v1/videos/generations": + postCount++ + var payload map[string]interface{} + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Fatalf("decode request body: %v", err) + } + if payload["model"] != "grok-imagine-video" || payload["duration"] != float64(8) { + t.Fatalf("unexpected request body: %#v", payload) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"request_id":"video_req_1"}`)) + case r.Method == http.MethodGet && r.URL.Path == "/v1/videos/video_req_1": + pollCount++ + w.Header().Set("Content-Type", "application/json") + if pollCount == 1 { + w.Header().Set("Retry-After", "0") + w.WriteHeader(http.StatusAccepted) + _, _ = w.Write([]byte(`{"status":"pending","progress":50}`)) + return + } + _, _ = w.Write([]byte(`{ + "status":"done", + "video":{ + "url":"https://example.com/generated.mp4", + "respect_moderation":true, + "file_output":{"filename":"generated.mp4"} + }, + "usage":{"cost_in_usd_ticks":27} + }`)) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + output, err := newTestClient().Generate(context.Background(), RouteConfig{ + Protocol: AdapterXAIVideo, + BaseURL: server.URL + "/v1", + APIKey: "xai-key", + ReadTimeoutMS: 5000, + UpstreamModel: "grok-imagine-video", + }, GenerateInput{ + Messages: []Message{{Role: "user", Content: "A cinematic orbit"}}, + Options: map[string]interface{}{"duration": 8}, + }) + if err != nil { + t.Fatalf("generate xAI video: %v", err) + } + if postCount != 1 || pollCount != 2 { + t.Fatalf("expected one submission and two polls, got post=%d poll=%d", postCount, pollCount) + } + if output.ResponseID != "video_req_1" || len(output.GeneratedVideos) != 1 { + t.Fatalf("unexpected xAI video output: %#v", output) + } + video := output.GeneratedVideos[0] + if video.URL != "https://example.com/generated.mp4" || video.MIMEType != "video/mp4" || video.FileName != "generated.mp4" { + t.Fatalf("unexpected generated video: %#v", video) + } + if !strings.Contains(output.Usage.RawUsageJSON, `"cost_in_usd_ticks":27`) { + t.Fatalf("expected raw usage JSON, got %#v", output.Usage) + } + if output.Debug == nil || output.Debug.Request.Path != "/v1/videos/video_req_1" { + t.Fatalf("expected final poll debug snapshot, got %#v", output.Debug) + } +} + +func TestGenerateXAIVideoMarksFailedTaskAsAccepted(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + if r.Method == http.MethodPost { + _, _ = w.Write([]byte(`{"request_id":"video_req_failed"}`)) + return + } + _, _ = w.Write([]byte(`{ + "status":"failed", + "error":{"code":"content_policy_violation","message":"request rejected"} + }`)) + })) + defer server.Close() + + _, err := newTestClient().Generate(context.Background(), RouteConfig{ + Protocol: AdapterXAIVideo, + BaseURL: server.URL + "/v1", + ReadTimeoutMS: 5000, + UpstreamModel: "grok-imagine-video", + }, GenerateInput{Messages: []Message{{Role: "user", Content: "Animate this"}}}) + if err == nil || !RequestWasAccepted(err) { + t.Fatalf("expected accepted request error, got %v", err) + } + if !strings.Contains(err.Error(), "content_policy_violation: request rejected") { + t.Fatalf("expected upstream failure details, got %v", err) + } +} + +func TestParseXAIVideoResultRejectsModeratedOutput(t *testing.T) { + _, _, err := parseXAIVideoResult([]byte(`{ + "status":"done", + "video":{"url":"https://example.com/blocked.mp4","respect_moderation":false} + }`), "video_req_1") + if err == nil || !strings.Contains(err.Error(), "content moderation") { + t.Fatalf("expected moderation error, got %v", err) + } +} diff --git a/backend/internal/infra/mediaartifact/client.go b/backend/internal/infra/mediaartifact/client.go index f97cfe01d..8d0f93c10 100644 --- a/backend/internal/infra/mediaartifact/client.go +++ b/backend/internal/infra/mediaartifact/client.go @@ -99,6 +99,7 @@ func (c *Client) DownloadImage(ctx context.Context, sourceURL string, trustedPro func (c *Client) DownloadVideo(ctx context.Context, sourceURL string, trustedProviderEndpoint string, apiKey string, maxBytes int64) ([]byte, string, error) { resolvedMIME := "" headers := map[string]string(nil) + providerBearerToken := strings.TrimSpace(apiKey) downloadURL := strings.TrimSpace(sourceURL) metadataURL, geminiDownloadURL, geminiFile := geminiGeneratedFileURLs(downloadURL) if geminiFile { @@ -112,16 +113,18 @@ func (c *Client) DownloadVideo(ctx context.Context, sourceURL string, trustedPro } downloadURL = geminiDownloadURL headers = map[string]string{geminiAPIKeyHeader: strings.TrimSpace(apiKey)} + providerBearerToken = "" } result, err := c.download(ctx, downloadRequest{ - url: downloadURL, - trustedEndpoint: trustedProviderEndpoint, - headers: headers, - maxBytes: maxBytes, - timeout: videoDownloadTimeout, - expectedMIMEPrefix: "video/", - failureLabel: "download generated video", + url: downloadURL, + trustedEndpoint: trustedProviderEndpoint, + headers: headers, + providerBearerToken: providerBearerToken, + maxBytes: maxBytes, + timeout: videoDownloadTimeout, + expectedMIMEPrefix: "video/", + failureLabel: "download generated video", }) if err != nil { return nil, "", err @@ -133,13 +136,14 @@ func (c *Client) DownloadVideo(ctx context.Context, sourceURL string, trustedPro } type downloadRequest struct { - url string - trustedEndpoint string - headers map[string]string - maxBytes int64 - timeout time.Duration - expectedMIMEPrefix string - failureLabel string + url string + trustedEndpoint string + headers map[string]string + providerBearerToken string + maxBytes int64 + timeout time.Duration + expectedMIMEPrefix string + failureLabel string } // download 统一执行带超时、状态码和响应大小边界的媒体下载。 @@ -166,6 +170,11 @@ func (c *Client) download(ctx context.Context, input downloadRequest) (downloadR for key, value := range input.headers { request.Header.Set(key, value) } + // 仅在制品 URL 与管理员配置的 Provider endpoint 明确同源时携带 Key。 + // 跨 origin 制品和后续跨 origin 重定向都不得获得 Provider 凭据。 + if trustedEndpoint != "" && sameArtifactOrigin(input.url, input.trustedEndpoint) && strings.TrimSpace(input.providerBearerToken) != "" { + request.Header.Set("Authorization", "Bearer "+strings.TrimSpace(input.providerBearerToken)) + } response, err := c.httpClients.Do(request, trustedEndpoint, "") if err != nil { return downloadResult{}, sanitizeRequestError(input.failureLabel, err) @@ -382,9 +391,16 @@ func stripCredentialOnCrossOriginRedirect(request *http.Request, via []*http.Req return nil } request.Header.Del(geminiAPIKeyHeader) + request.Header.Del("Authorization") return nil } +func sameArtifactOrigin(sourceURL string, providerEndpoint string) bool { + sourceOrigin, sourceErr := security.HTTPOrigin(strings.TrimSpace(sourceURL)) + providerOrigin, providerErr := security.HTTPOrigin(strings.TrimSpace(providerEndpoint)) + return sourceErr == nil && providerErr == nil && sourceOrigin == providerOrigin +} + func mediaArtifactRedirectPolicy(strictPolicy security.OutboundPolicy, trustedOrigin string) func(*http.Request, []*http.Request) error { return func(request *http.Request, via []*http.Request) error { if err := stripCredentialOnCrossOriginRedirect(request, via); err != nil { diff --git a/backend/internal/infra/mediaartifact/client_test.go b/backend/internal/infra/mediaartifact/client_test.go index 7ff111a4e..6f318ce18 100644 --- a/backend/internal/infra/mediaartifact/client_test.go +++ b/backend/internal/infra/mediaartifact/client_test.go @@ -154,6 +154,46 @@ func TestDownloadVideoRequiresGeminiAPIKeyBeforeRequest(t *testing.T) { } } +func TestDownloadVideoUsesBearerTokenForSameOriginProviderArtifact(t *testing.T) { + client := testClient(roundTripFunc(func(request *http.Request) (*http.Response, error) { + if got := request.Header.Get("Authorization"); got != "Bearer provider-key" { + t.Fatalf("expected same-origin provider authorization, got %q", got) + } + return response(http.StatusOK, "video/mp4", []byte("video-bytes")), nil + })) + + _, _, err := client.DownloadVideo( + t.Context(), + "https://provider.example.test/v1/artifacts/video.mp4", + "https://provider.example.test/v1", + "provider-key", + 1024, + ) + if err != nil { + t.Fatalf("download same-origin provider video: %v", err) + } +} + +func TestDownloadVideoDoesNotSendBearerTokenToCrossOriginArtifact(t *testing.T) { + client := testClient(roundTripFunc(func(request *http.Request) (*http.Response, error) { + if got := request.Header.Get("Authorization"); got != "" { + t.Fatalf("provider authorization leaked to cross-origin artifact: %q", got) + } + return response(http.StatusOK, "video/mp4", []byte("video-bytes")), nil + })) + + _, _, err := client.DownloadVideo( + t.Context(), + "https://cdn.example.test/generated/video.mp4", + "https://provider.example.test/v1", + "provider-key", + 1024, + ) + if err != nil { + t.Fatalf("download cross-origin provider video: %v", err) + } +} + func TestGeminiMetadataErrorDoesNotExposeResponseBody(t *testing.T) { client := testClient(roundTripFunc(func(*http.Request) (*http.Response, error) { return response(http.StatusBadGateway, "application/json", []byte(`{"error":"token=secret user-content"}`)), nil @@ -174,7 +214,7 @@ func TestGeminiMetadataErrorDoesNotExposeResponseBody(t *testing.T) { } } -func TestRedirectPolicyStripsGeminiKeyAcrossOrigins(t *testing.T) { +func TestRedirectPolicyStripsCredentialsAcrossOrigins(t *testing.T) { originalURL, err := url.Parse("https://generativelanguage.googleapis.com/v1beta/files/video_123:download") if err != nil { t.Fatal(err) @@ -185,6 +225,7 @@ func TestRedirectPolicyStripsGeminiKeyAcrossOrigins(t *testing.T) { } original := &http.Request{URL: originalURL, Header: make(http.Header)} original.Header.Set(geminiAPIKeyHeader, "secret") + original.Header.Set("Authorization", "Bearer provider-secret") redirect := &http.Request{URL: redirectURL, Header: original.Header.Clone()} if err = stripCredentialOnCrossOriginRedirect(redirect, []*http.Request{original}); err != nil { t.Fatalf("check redirect: %v", err) @@ -192,6 +233,9 @@ func TestRedirectPolicyStripsGeminiKeyAcrossOrigins(t *testing.T) { if redirect.Header.Get(geminiAPIKeyHeader) != "" { t.Fatal("Gemini API key leaked to a different origin") } + if redirect.Header.Get("Authorization") != "" { + t.Fatal("provider authorization leaked to a different origin") + } sameOriginRedirect := &http.Request{URL: originalURL, Header: original.Header.Clone()} if err = stripCredentialOnCrossOriginRedirect(sameOriginRedirect, []*http.Request{original}); err != nil { @@ -200,6 +244,9 @@ func TestRedirectPolicyStripsGeminiKeyAcrossOrigins(t *testing.T) { if sameOriginRedirect.Header.Get(geminiAPIKeyHeader) != "secret" { t.Fatal("same-origin redirect unexpectedly removed Gemini API key") } + if sameOriginRedirect.Header.Get("Authorization") != "Bearer provider-secret" { + t.Fatal("same-origin redirect unexpectedly removed provider authorization") + } canonicalURL, err := url.Parse("https://GENERATIVELANGUAGE.googleapis.com:443/v1beta/files/video_123:download") if err != nil { @@ -212,6 +259,9 @@ func TestRedirectPolicyStripsGeminiKeyAcrossOrigins(t *testing.T) { if canonicalRedirect.Header.Get(geminiAPIKeyHeader) != "secret" { t.Fatal("canonical same-origin redirect unexpectedly removed Gemini API key") } + if canonicalRedirect.Header.Get("Authorization") != "Bearer provider-secret" { + t.Fatal("canonical same-origin redirect unexpectedly removed provider authorization") + } } func TestRedirectPolicyStopsAfterLimit(t *testing.T) { diff --git a/frontend/features/admin/api/llm.types.ts b/frontend/features/admin/api/llm.types.ts index 709e96fd6..29472b8a5 100644 --- a/frontend/features/admin/api/llm.types.ts +++ b/frontend/features/admin/api/llm.types.ts @@ -56,7 +56,8 @@ export type AdminLLMAdapter = | "gemini_interactions" | "xai_responses" | "xai_image" - | "xai_image_edits"; + | "xai_image_edits" + | "xai_video"; export type AdminLLMModelVendor = string; export type AdminLLMCompatible = | "openai" diff --git a/frontend/features/admin/components/sections/conversation/admin-conversation.tsx b/frontend/features/admin/components/sections/conversation/admin-conversation.tsx index d48c93f62..fede411d2 100644 --- a/frontend/features/admin/components/sections/conversation/admin-conversation.tsx +++ b/frontend/features/admin/components/sections/conversation/admin-conversation.tsx @@ -286,6 +286,11 @@ generationConfig.safetySettings.threshold`} "resolution", "response_format" ], + "xai_video": [ + "aspect_ratio", + "duration", + "resolution" + ], "openai_chat_completions": [ "service_tier", "thinking.type" @@ -332,7 +337,7 @@ generationConfig.safetySettings.threshold`}
{t("guide.protocolDescription")}
{item}
))}