diff --git a/apps/web/src/pages/PlaygroundPage.models.test.ts b/apps/web/src/pages/PlaygroundPage.models.test.ts new file mode 100644 index 0000000..dd52e6b --- /dev/null +++ b/apps/web/src/pages/PlaygroundPage.models.test.ts @@ -0,0 +1,48 @@ +import type { PlatformModel } from '@easyai-ai-gateway/contracts'; +import { describe, expect, it } from 'vitest'; +import { filterModelsForMode } from './PlaygroundPage'; + +function model(id: string, modelType: string[]) { + return { id, modelType } as PlatformModel; +} + +function modelIds(models: PlatformModel[]) { + return models.map((item) => item.id); +} + +describe('playground model filtering', () => { + const models = [ + model('image-generation', ['image_generate']), + model('legacy-image', ['image']), + model('image-edit', ['image_edit']), + model('image-to-video', ['video_generate', 'image_to_video']), + model('text-to-video', ['text_to_video']), + ]; + + it('only shows image generation models when no reference image is present', () => { + expect(modelIds(filterModelsForMode(models, 'image', false, 'text_to_video'))).toEqual([ + 'image-generation', + 'legacy-image', + ]); + }); + + it('only shows image editing models when a reference image is present', () => { + expect(modelIds(filterModelsForMode(models, 'image', true, 'text_to_video'))).toEqual([ + 'legacy-image', + 'image-edit', + ]); + }); + + it('does not fall back to image-to-video models when no image model is available', () => { + const videoOnlyModels = [ + model('kling-3-turbo', ['video_generate', 'image_to_video']), + model('kling-1-5', ['image_to_video']), + ]; + + expect(filterModelsForMode(videoOnlyModels, 'image', false, 'text_to_video')).toEqual([]); + }); + + it('keeps image-to-video models available for the matching video mode', () => { + expect(modelIds(filterModelsForMode(models, 'video', false, 'first_last_frame'))).toContain('image-to-video'); + }); +}); diff --git a/apps/web/src/pages/PlaygroundPage.tsx b/apps/web/src/pages/PlaygroundPage.tsx index 38daa15..e671aaf 100644 --- a/apps/web/src/pages/PlaygroundPage.tsx +++ b/apps/web/src/pages/PlaygroundPage.tsx @@ -810,7 +810,7 @@ function Composer(props: { {props.mode !== 'chat' && props.mediaSettings && props.onMediaSettingsChange && ( = { first_last_frame: ['video_first_last_frame', 'image_to_video', 'video_generate'], omni_reference: ['omni_video', 'video_reference', 'video_generate'], text_to_video: ['text_to_video', 'video_generate'], }; - return filterWithFallback(models, [...videoTypesByMode[videoMode], 'video']); + return filterModelsByType(models, [...videoTypesByMode[videoMode], 'video']); } -function filterWithFallback(models: PlatformModel[], modelTypes: string[]) { - const exact = models.filter((model) => model.modelType.some((type) => modelTypes.includes(type))); - return exact.length ? exact : models.filter((model) => modelTypes.some((type) => model.modelType.some((modelType) => modelType.includes(type) || type.includes(modelType)))); +function filterModelsByType(models: PlatformModel[], modelTypes: string[]) { + const acceptedTypes = new Set(modelTypes); + return models.filter((model) => model.modelType.some((type) => acceptedTypes.has(type))); } function buildModelOptions(models: PlatformModel[]): ModelOption[] {