diff --git a/src/components/__tests__/GenerateAudioNode.test.tsx b/src/components/__tests__/GenerateAudioNode.test.tsx index 50abb4a0..06810c04 100644 --- a/src/components/__tests__/GenerateAudioNode.test.tsx +++ b/src/components/__tests__/GenerateAudioNode.test.tsx @@ -36,6 +36,10 @@ vi.mock("@/store/modelStore", () => ({ multiAngle: [], highDefinition: [], outpainting: [], + inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, all: [ ...mockPopiserverModelLists.image.models, diff --git a/src/components/__tests__/GenerateImageNode.test.tsx b/src/components/__tests__/GenerateImageNode.test.tsx index 16dfdf1e..34f23066 100644 --- a/src/components/__tests__/GenerateImageNode.test.tsx +++ b/src/components/__tests__/GenerateImageNode.test.tsx @@ -51,6 +51,10 @@ vi.mock("@/store/modelStore", () => ({ multiAngle: [], highDefinition: [], outpainting: [], + inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, all: [ ...mockPopiserverModelLists.image.models, diff --git a/src/components/__tests__/GenerateVideoNode.test.tsx b/src/components/__tests__/GenerateVideoNode.test.tsx index bbf9bffa..76d5af44 100644 --- a/src/components/__tests__/GenerateVideoNode.test.tsx +++ b/src/components/__tests__/GenerateVideoNode.test.tsx @@ -55,6 +55,10 @@ vi.mock("@/store/modelStore", () => ({ multiAngle: [], highDefinition: [], outpainting: [], + inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, all: [ ...mockPopiserverModelLists.image.models, @@ -1138,6 +1142,55 @@ describe("GenerateVideoNode", () => { String(call[0]).startsWith("/api/models/99") )).toBe(false); }); + + it("does not auto-switch a high-definition derived node to a default video model", async () => { + // 高清派生节点用的是 206 增强模型(不在 byKind.video 生成模型列表里)。 + // 自动选型必须跳过派生视频节点,否则会把模型替换成默认视频模型并多查一次其详情。 + mockFetch.mockImplementation((url) => { + const requestUrl = String(url); + if (requestUrl.startsWith("/api/models?")) { + return Promise.resolve({ + ok: true, + json: () => Promise.resolve({ + models: [ + { id: "32", name: "First Popi Video", provider: "popiserver", capabilities: ["text-to-video"] }, + ], + }), + }); + } + return Promise.resolve({ + ok: true, + json: () => Promise.resolve({ parameters: [], inputs: [] }), + }); + }); + mockPopiserverModelLists.video.models = [ + { id: "32", name: "First Popi Video", provider: "popiserver", capabilities: ["text-to-video"], description: null }, + ]; + + render( + + )} /> + + ); + + // 给 effect 一个执行机会 + await new Promise((resolve) => setTimeout(resolve, 0)); + + // 不应把节点模型替换成默认视频模型(id 32) + expect(mockUpdateNodeData).not.toHaveBeenCalledWith( + "test-node-1", + expect.objectContaining({ + selectedModel: expect.objectContaining({ modelId: "32" }), + }) + ); + // 不应多查默认模型 32 的详情 + expect(mockFetch.mock.calls.some((call) => + String(call[0]).startsWith("/api/models/32") + )).toBe(false); + }); }); describe("Fetch Models on Mount", () => { diff --git a/src/components/__tests__/GenerationComposer.test.tsx b/src/components/__tests__/GenerationComposer.test.tsx index 72b355a7..2788c042 100644 --- a/src/components/__tests__/GenerationComposer.test.tsx +++ b/src/components/__tests__/GenerationComposer.test.tsx @@ -770,6 +770,10 @@ describe("GenerationComposer", () => { multiAngle: [], highDefinition: [], outpainting: [], + inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, detailsById: {}, detailLoadingById: {}, diff --git a/src/components/__tests__/LLMGenerateNode.test.tsx b/src/components/__tests__/LLMGenerateNode.test.tsx index 94c7382a..f7f07b6e 100644 --- a/src/components/__tests__/LLMGenerateNode.test.tsx +++ b/src/components/__tests__/LLMGenerateNode.test.tsx @@ -32,6 +32,10 @@ vi.mock("@/store/modelStore", () => ({ multiAngle: [], highDefinition: [], outpainting: [], + inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, detailsById: {}, detailLoadingById: {}, diff --git a/src/components/__tests__/ModelSearchDialog.test.tsx b/src/components/__tests__/ModelSearchDialog.test.tsx index e0b41855..4f3c650c 100644 --- a/src/components/__tests__/ModelSearchDialog.test.tsx +++ b/src/components/__tests__/ModelSearchDialog.test.tsx @@ -60,6 +60,10 @@ vi.mock("@/store/modelStore", () => ({ multiAngle: mockPopiModels.filter((model) => model.metadata?.subType === 113), highDefinition: mockPopiModels.filter((model) => model.metadata?.subType === 107), outpainting: mockPopiModels.filter((model) => model.metadata?.subType === 112), + inpainting: mockPopiModels.filter((model) => model.metadata?.subType === 110), + videoEnhance: mockPopiModels.filter((model) => model.metadata?.subType === 206), + videoSmartErase: mockPopiModels.filter((model) => model.metadata?.subType === 207), + videoBoxErase: mockPopiModels.filter((model) => model.metadata?.subType === 208), }, initializePopiModels: vi.fn(), refreshPopiModels: mockRefreshPopiModels, diff --git a/src/components/canvas/ActiveVideoEditHost.tsx b/src/components/canvas/ActiveVideoEditHost.tsx index 11c0f567..96d0d647 100644 --- a/src/components/canvas/ActiveVideoEditHost.tsx +++ b/src/components/canvas/ActiveVideoEditHost.tsx @@ -24,7 +24,8 @@ export function ActiveVideoEditHost() { const targetNodeId = isVideoEditTool ? activeNodeId : null; const frame = useNodeCanvasFrame(targetNodeId); - const videoModels = useModelStore((state) => state.byKind.video); + const videoSmartEraseModels = useModelStore((state) => state.byKind.videoSmartErase); + const videoBoxEraseModels = useModelStore((state) => state.byKind.videoBoxErase); const videoSource = frame ? getNodeVideoSource(frame.node) : null; @@ -50,14 +51,14 @@ export function ActiveVideoEditHost() { ) : ( )} diff --git a/src/components/composer/NodeHighDefinitionPanel.tsx b/src/components/composer/NodeHighDefinitionPanel.tsx index c24b520f..bca97111 100644 --- a/src/components/composer/NodeHighDefinitionPanel.tsx +++ b/src/components/composer/NodeHighDefinitionPanel.tsx @@ -20,7 +20,6 @@ import { buildHighDefinitionParameters, chooseHighDefinitionModel, HIGH_DEFINITI import { chooseHighDefinitionVideoModel, isHighDefinitionDerivedVideoNode, - isHighDefinitionVideoModel, } from "@/utils/highDefinitionVideoNodes"; import { toSelectedModel } from "@/utils/selectedModel"; @@ -74,11 +73,7 @@ export function NodeHighDefinitionPanel({ node }: NodeHighDefinitionPanelProps) const isRunning = useWorkflowStore((state) => state.runningNodeIds.has(node.id)); const isVideoHighDefinitionNode = isHighDefinitionDerivedVideoNode(node.data); const highDefinitionModels = useModelStore((state) => state.byKind.highDefinition); - const videoModels = useModelStore((state) => state.byKind.video); - const videoEnhanceModels = useMemo( - () => videoModels.filter(isHighDefinitionVideoModel), - [videoModels] - ); + const videoEnhanceModels = useModelStore((state) => state.byKind.videoEnhance); const fetchModelDetail = useModelStore((state) => state.fetchModelDetail); const data = node.data as SmartImageNodeData | SmartVideoNodeData; const hydratedKeyRef = useRef(null); diff --git a/src/components/nodes/GenerateVideoNode.tsx b/src/components/nodes/GenerateVideoNode.tsx index bb651277..6df04247 100644 --- a/src/components/nodes/GenerateVideoNode.tsx +++ b/src/components/nodes/GenerateVideoNode.tsx @@ -33,7 +33,7 @@ import { asMediaDimensions, getAutoMediaElementClassName, getAutoMediaFrameClass import { CustomVideoPlayer } from "@/components/media/CustomVideoPlayer"; import { createVideoFrameSmartImageNode } from "@/utils/videoFrameNode"; import { chooseHighDefinitionVideoModel, createHighDefinitionVideoNode } from "@/utils/highDefinitionVideoNodes"; -import { isSmartVideoAutoTaskNode } from "@/utils/smartMediaMode"; +import { getDerivedMediaInfo } from "@/utils/derivedMediaNodes"; import { getActiveNodeTool, useNodeToolStore } from "@/store/nodeToolStore"; import { createVideoEraseActionPanel } from "@/components/media/MediaEditActionDropdown"; import { VideoEraseMenuIcon } from "@/components/media/MediaToolbarIcons"; @@ -68,6 +68,7 @@ export function GenerateVideoNodeView({ id, data, selected, suppressEmptyStateDe const updateMediaNodeData = useWorkflowStore((state) => state.updateMediaNodeData); const updateVideoPreference = useGenerationPreferenceStore((state) => state.updateVideoPreference); const popiserverVideoModels = useModelStore((state) => state.byKind.video); + const popiserverVideoEnhanceModels = useModelStore((state) => state.byKind.videoEnhance); const [isBrowseDialogOpen, setIsBrowseDialogOpen] = useState(false); const [isLoadingCarouselVideo, setIsLoadingCarouselVideo] = useState(false); const [isPreviewOpen, setIsPreviewOpen] = useState(false); @@ -115,9 +116,9 @@ export function GenerateVideoNodeView({ id, data, selected, suppressEmptyStateDe useEffect(() => { if (modelOptions.length === 0) return; - // 视频派生节点使用内置模型或不需要生成模型(裁剪为本地处理),不参与 - // 自动选型,否则会被替换成默认视频生成模型、丢失派生语义。 - if (isSmartVideoAutoTaskNode(nodeData)) return; + // 视频派生节点(裁剪/擦除/高清)使用各自的内置或专属模型,不参与自动选型, + // 否则会被替换成默认视频生成模型、丢失派生语义并多触发一次默认模型详情请求。 + if (getDerivedMediaInfo(nodeData)?.field === "derivedVideo") return; const selectedModelExists = modelOptions.some((model) => model.id === nodeConfig.selectedModel?.modelId); if (nodeConfig.selectedModel?.provider === currentProvider && nodeConfig.selectedModel.modelId && selectedModelExists) return; const selectedModel = chooseModelFromList(modelOptions, nodeConfig.selectedModel); @@ -363,12 +364,12 @@ export function GenerateVideoNodeView({ id, data, selected, suppressEmptyStateDe y: currentPosition.y, }, dimensions: asStrictMediaDimensions(nodeData.dimensions), - selectedModel: chooseHighDefinitionVideoModel(popiserverVideoModels), + selectedModel: chooseHighDefinitionVideoModel(popiserverVideoEnhanceModels), addNode, onConnect, selectSingleNode, }); - }, [addNode, getNodeById, id, nodeData.dimensions, nodeData.outputVideo, onConnect, popiserverVideoModels, selectSingleNode]); + }, [addNode, getNodeById, id, nodeData.dimensions, nodeData.outputVideo, onConnect, popiserverVideoEnhanceModels, selectSingleNode]); const switchNodeTool = useNodeToolStore((state) => state.switchTool); diff --git a/src/components/nodes/VideoInputNode.tsx b/src/components/nodes/VideoInputNode.tsx index 1bddab8d..2605d70a 100644 --- a/src/components/nodes/VideoInputNode.tsx +++ b/src/components/nodes/VideoInputNode.tsx @@ -87,7 +87,7 @@ export function VideoInputNodeView({ id, data, selected }: VideoInputNodeViewPro const getNodeById = useWorkflowStore((state) => state.getNodeById); const selectSingleNode = useWorkflowStore((state) => state.selectSingleNode); const regenerateNode = useWorkflowStore((state) => state.regenerateNode); - const popiserverVideoModels = useModelStore((state) => state.byKind.video); + const popiserverVideoEnhanceModels = useModelStore((state) => state.byKind.videoEnhance); const fileInputRef = useRef(null); const [isPreviewOpen, setIsPreviewOpen] = useState(false); const [isAssetPickerOpen, setIsAssetPickerOpen] = useState(false); @@ -280,12 +280,12 @@ export function VideoInputNodeView({ id, data, selected }: VideoInputNodeViewPro y: currentPosition.y, }, dimensions: nodeData.dimensions, - selectedModel: chooseHighDefinitionVideoModel(popiserverVideoModels), + selectedModel: chooseHighDefinitionVideoModel(popiserverVideoEnhanceModels), addNode, onConnect, selectSingleNode, }); - }, [addNode, getNodeById, id, nodeData.dimensions, nodeData.video, onConnect, popiserverVideoModels, selectSingleNode]); + }, [addNode, getNodeById, id, nodeData.dimensions, nodeData.video, onConnect, popiserverVideoEnhanceModels, selectSingleNode]); const switchNodeTool = useNodeToolStore((state) => state.switchTool); diff --git a/src/store/__tests__/modelStore.test.ts b/src/store/__tests__/modelStore.test.ts index 106ab934..187295f3 100644 --- a/src/store/__tests__/modelStore.test.ts +++ b/src/store/__tests__/modelStore.test.ts @@ -1,7 +1,11 @@ import { beforeEach, describe, expect, it } from "vitest"; import { useModelStore } from "@/store/modelStore"; import type { ProviderModel } from "@/lib/providers/types"; -import { POPI_VIDEO_SUBTYPE_ENHANCE } from "@/constants/generationTask"; +import { + POPI_VIDEO_SUBTYPE_ENHANCE, + POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE, + POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO, +} from "@/constants/generationTask"; function model(id: string, subType: number): ProviderModel { return { @@ -32,6 +36,9 @@ describe("modelStore", () => { highDefinition: [], outpainting: [], inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }, detailsById: {}, detailLoadingById: {}, @@ -75,7 +82,7 @@ describe("modelStore", () => { expect(byKind.image).toEqual([regularImage]); }); - it("groups subtype 206 models into the video category", () => { + it("groups subtype 206/207/208 models into their dedicated video tool categories", () => { const regularImage = model("regular-image", 103); const videoEnhance = { ...model("video-enhance", POPI_VIDEO_SUBTYPE_ENHANCE), @@ -86,11 +93,32 @@ describe("modelStore", () => { subTypes: [POPI_VIDEO_SUBTYPE_ENHANCE], }, } satisfies ProviderModel; + const videoSmartErase = { + ...model("video-smart-erase", POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE), + capabilities: ["image-to-video"], + metadata: { + type: 2, + subType: POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE, + subTypes: [POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE], + }, + } satisfies ProviderModel; + const videoBoxErase = { + ...model("video-box-erase", POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO), + capabilities: ["image-to-video"], + metadata: { + type: 2, + subType: POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO, + subTypes: [POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO], + }, + } satisfies ProviderModel; - useModelStore.getState().setPopiModels([regularImage, videoEnhance]); + useModelStore.getState().setPopiModels([regularImage, videoEnhance, videoSmartErase, videoBoxErase]); const byKind = useModelStore.getState().byKind; - expect(byKind.video).toEqual([videoEnhance]); + expect(byKind.videoEnhance).toEqual([videoEnhance]); + expect(byKind.videoSmartErase).toEqual([videoSmartErase]); + expect(byKind.videoBoxErase).toEqual([videoBoxErase]); + expect(byKind.video).toEqual([]); expect(byKind.image).toEqual([regularImage]); }); }); diff --git a/src/store/modelStore.ts b/src/store/modelStore.ts index 534bbf74..fd20e998 100644 --- a/src/store/modelStore.ts +++ b/src/store/modelStore.ts @@ -1,12 +1,32 @@ import { create } from "zustand"; import type { ModelInput, ModelParameter, ProviderModel, ModelCapability } from "@/lib/providers/types"; +import { + POPI_VIDEO_SUBTYPE_ENHANCE, + POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE, + POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO, + POPI_VIDEO_SUBTYPE_FIRST_LAST_FRAME, + POPI_VIDEO_SUBTYPE_IMAGE_TO_VIDEO, + POPI_VIDEO_SUBTYPE_OMNI_REFERENCE, +} from "@/constants/generationTask"; import { getPopiCapabilitiesForKind, getPopiKindForSubType, type PopiModelMediaKind, } from "@/utils/popiModelClassification"; -export type PopiModelKind = "image" | "video" | "audio" | "3d" | "llm" | "multiAngle" | "highDefinition" | "outpainting" | "inpainting"; +export type PopiModelKind = + | "image" + | "video" + | "audio" + | "3d" + | "llm" + | "multiAngle" + | "highDefinition" + | "outpainting" + | "inpainting" + | "videoEnhance" + | "videoSmartErase" + | "videoBoxErase"; export type PopiModelExtensionFields = Record | unknown[]; interface ModelListState { @@ -55,6 +75,9 @@ const emptyByKind: Record = { highDefinition: [], outpainting: [], inpainting: [], + videoEnhance: [], + videoSmartErase: [], + videoBoxErase: [], }; let popiModelsRequest: Promise | null = null; @@ -78,10 +101,18 @@ function hasPopiKind(model: ProviderModel, kind: PopiModelMediaKind): boolean { return hasCapability(model, getPopiCapabilitiesForKind(kind)); } +function hasAnySubType(model: ProviderModel, subTypes: number[]): boolean { + return getMetadataSubTypes(model).some((subType) => subTypes.includes(subType)); +} + function groupPopiModels(models: ProviderModel[]): Record { return { image: models.filter((model) => getMetadataSubTypes(model).some((subType) => subType === 102 || subType === 103)), - video: models.filter((model) => getMetadataSubTypes(model).some((subType) => subType === 202 || subType === 203 || subType === 204)), + video: models.filter((model) => hasAnySubType(model, [ + POPI_VIDEO_SUBTYPE_IMAGE_TO_VIDEO, + POPI_VIDEO_SUBTYPE_OMNI_REFERENCE, + POPI_VIDEO_SUBTYPE_FIRST_LAST_FRAME, + ])), audio: models.filter((model) => hasPopiKind(model, "audio")), "3d": models.filter((model) => hasCapability(model, ["text-to-3d", "image-to-3d"])), llm: models.filter((model) => hasPopiKind(model, "llm") || Number(model.metadata?.type) === 5), @@ -89,6 +120,9 @@ function groupPopiModels(models: ProviderModel[]): Record getMetadataSubTypes(model).includes(107)), outpainting: models.filter((model) => getMetadataSubTypes(model).includes(112)), inpainting: models.filter((model) => getMetadataSubTypes(model).includes(110)), + videoEnhance: models.filter((model) => hasAnySubType(model, [POPI_VIDEO_SUBTYPE_ENHANCE])), + videoSmartErase: models.filter((model) => hasAnySubType(model, [POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE])), + videoBoxErase: models.filter((model) => hasAnySubType(model, [POPI_VIDEO_SUBTYPE_ERASE_SUBTITLE_PRO])), }; } diff --git a/src/utils/highDefinitionVideoNodes.ts b/src/utils/highDefinitionVideoNodes.ts index b201c742..b33d39e7 100644 --- a/src/utils/highDefinitionVideoNodes.ts +++ b/src/utils/highDefinitionVideoNodes.ts @@ -12,6 +12,7 @@ export const HIGH_DEFINITION_VIDEO_FALLBACK_MODEL: SelectedModel = { modelId: "", displayName: "画质增强", capabilities: ["image-to-video"], + isSupportVideos: true, type: 2, subType: POPI_VIDEO_SUBTYPE_ENHANCE, subTypes: [POPI_VIDEO_SUBTYPE_ENHANCE],