Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 11 additions & 3 deletions backend/internal/application/channel/model_catalog.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ const (
protocolGeminiInteractions = llm.AdapterGeminiInteractions
protocolXAIImage = llm.AdapterXAIImage
protocolXAIImageEdits = llm.AdapterXAIImageEdits
protocolXAIVideo = llm.AdapterXAIVideo
)

var protocolDefaultKindOrder = []string{
Expand Down Expand Up @@ -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{
Expand Down Expand Up @@ -173,7 +175,8 @@ func isKnownProtocol(raw string) bool {
protocolGoogleImageGeneration,
protocolGeminiInteractions,
protocolXAIImage,
protocolXAIImageEdits:
protocolXAIImageEdits,
protocolXAIVideo:
return true
default:
return false
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"]`
Expand Down Expand Up @@ -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")
}
11 changes: 6 additions & 5 deletions backend/internal/application/channel/model_catalog_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}

Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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")
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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": {},
Expand Down
12 changes: 8 additions & 4 deletions backend/internal/application/settings/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -198,12 +198,16 @@ func isLegacyDefaultModelOptionAllowedPaths(value string) bool {
if err := json.Unmarshal([]byte(strings.TrimSpace(value)), &current); 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 {
Expand Down
27 changes: 27 additions & 0 deletions backend/internal/application/settings/service_seed_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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{
Expand Down
5 changes: 5 additions & 0 deletions backend/internal/infra/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,11 @@ func DefaultModelOptionAllowedPathsJSON() string {
"resolution",
"response_format"
],
"xai_video": [
"aspect_ratio",
"duration",
"resolution"
],
"gemini_generate_content": [
"generationConfig.temperature",
"generationConfig.topP",
Expand Down
15 changes: 12 additions & 3 deletions backend/internal/infra/llm/adapter.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -62,7 +63,8 @@ func IsKnownAdapter(raw string) bool {
AdapterGeminiInteractions,
AdapterXAIResponses,
AdapterXAIImage,
AdapterXAIImageEdits:
AdapterXAIImageEdits,
AdapterXAIVideo:
return true
default:
return false
Expand All @@ -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
Expand Down Expand Up @@ -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 返回协议对应的固定端点标识。
Expand All @@ -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:
Expand Down
15 changes: 15 additions & 0 deletions backend/internal/infra/llm/adapter_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
5 changes: 5 additions & 0 deletions backend/internal/infra/llm/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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},
Expand Down Expand Up @@ -1637,6 +1640,8 @@ func normalizeEndpoint(raw string) string {
return EndpointImageGenerations
case EndpointImageEdits:
return EndpointImageEdits
case EndpointVideoGenerations:
return EndpointVideoGenerations
case EndpointInteractions:
return EndpointInteractions
default:
Expand Down
6 changes: 6 additions & 0 deletions backend/internal/infra/llm/endpoint_url_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 2 additions & 0 deletions backend/internal/infra/llm/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Expand Down
Loading
Loading