feat: 生成增加停止功能

This commit is contained in:
soldier
2026-06-16 18:44:20 +08:00
parent 1b289f6a77
commit 73090da95a
6 changed files with 230 additions and 75 deletions
@@ -85,10 +85,18 @@ type CanvasHistoryEntry = Pick<CanvasClipboard, "nodes" | "connections"> & {
showImageInfo: boolean;
};
type CanvasGenerationRequest = {
targetNodeId: string;
originNodeId: string;
runningNodeId: string;
controller: AbortController;
};
const VIDEO_NODE_MAX_WIDTH = 420;
const VIDEO_NODE_MAX_HEIGHT = 420;
const CONNECTION_HANDLE_HIT_RADIUS = 40;
const CONNECTION_NODE_HIT_PADDING = 32;
const NODE_STATUS_IDLE = "idle" as const;
const NODE_STATUS_LOADING = "loading" as const;
const NODE_STATUS_SUCCESS = "success" as const;
const NODE_STATUS_ERROR = "error" as const;
@@ -208,7 +216,7 @@ function ConnectionCreateOption({ theme, icon, title, description, onClick }: {
}
function InfiniteCanvasPage() {
const { message } = App.useApp();
const { message, modal } = App.useApp();
const params = useParams<{ id: string }>();
const router = useRouter();
const projectId = params.id;
@@ -311,6 +319,7 @@ function InfiniteCanvasPage() {
const selectionBoxRef = useRef(selectionBox);
const agentCloseTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null);
const pendingConnectionCreateRef = useRef(pendingConnectionCreate);
const generationRequestsRef = useRef(new Map<string, CanvasGenerationRequest>());
const createHistoryEntry = useCallback(
(): CanvasHistoryEntry => ({
@@ -331,6 +340,52 @@ function InfiniteCanvasPage() {
[cleanupAssetImages],
);
const startGenerationRequest = useCallback((targetNodeId: string, originNodeId: string, runningId = originNodeId, controller = new AbortController()) => {
const previous = generationRequestsRef.current.get(targetNodeId);
if (previous?.controller !== controller) previous?.controller.abort();
generationRequestsRef.current.set(targetNodeId, { targetNodeId, originNodeId, runningNodeId: runningId, controller });
return controller;
}, []);
const finishGenerationRequest = useCallback((targetNodeId: string, controller: AbortController) => {
const request = generationRequestsRef.current.get(targetNodeId);
if (request?.controller === controller) generationRequestsRef.current.delete(targetNodeId);
}, []);
const stopGenerationByRunningId = useCallback((runningId: string) => {
const affectedNodeIds = new Set<string>();
generationRequestsRef.current.forEach((request) => {
if (request.runningNodeId !== runningId) return;
request.controller.abort();
generationRequestsRef.current.delete(request.targetNodeId);
affectedNodeIds.add(request.targetNodeId);
affectedNodeIds.add(request.originNodeId);
});
setRunningNodeId((current) => (current === runningId ? null : current));
if (!affectedNodeIds.size) return;
setNodes((prev) =>
prev.map((node) =>
affectedNodeIds.has(node.id) && node.metadata?.status === NODE_STATUS_LOADING
? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_IDLE, errorDetails: undefined } }
: node,
),
);
}, []);
const confirmStopGeneration = useCallback(
(nodeId: string) => {
modal.confirm({
title: "停止生成?",
content: "当前生成请求会被中断,已经生成完成的内容会保留。",
okText: "停止",
cancelText: "继续生成",
okButtonProps: { danger: true },
onOk: () => stopGenerationByRunningId(nodeId),
});
},
[modal, stopGenerationByRunningId],
);
useEffect(() => {
if (!hydrated) return;
setProjectLoaded(false);
@@ -1653,20 +1708,23 @@ function InfiniteCanvasPage() {
setSelectedNodeIds(new Set([childId]));
setSelectedConnectionId(null);
setDialogNodeId(childId);
const controller = startGenerationRequest(childId, node.id, childId);
try {
const image = await requestEdit(generationConfig, prompt, [source], { id: `${node.id}-mask`, name: "mask.png", type: "image/png", dataUrl: payload.maskDataUrl }).then((items) => items[0]);
const image = await requestEdit(generationConfig, prompt, [source], { id: `${node.id}-mask`, name: "mask.png", type: "image/png", dataUrl: payload.maskDataUrl }, { signal: controller.signal }).then((items) => items[0]);
const uploaded = await uploadImage(image.dataUrl);
const size = fitNodeSize(uploaded.width, uploaded.height, node.width, node.height);
setNodes((prev) => prev.map((item) => (item.id === childId ? { ...item, width: size.width, height: size.height, metadata: { ...item.metadata, ...imageMetadata(uploaded), prompt, ...generationMetadata } } : item)));
} catch (error) {
if (isGenerationCanceled(error)) return;
const errorDetails = error instanceof Error ? error.message : "局部修改失败";
message.error(errorDetails);
setNodes((prev) => prev.map((item) => (item.id === childId ? { ...item, metadata: { ...item.metadata, status: NODE_STATUS_ERROR, errorDetails } } : item)));
} finally {
finishGenerationRequest(childId, controller);
setRunningNodeId(null);
}
},
[effectiveConfig, isAiConfigReady, message, openConfigDialog],
[effectiveConfig, finishGenerationRequest, isAiConfigReady, message, openConfigDialog, startGenerationRequest],
);
const upscaleImageNode = useCallback(async (node: CanvasNodeData, params: CanvasImageUpscaleParams) => {
@@ -1726,21 +1784,24 @@ function InfiniteCanvasPage() {
setConnections((prev) => [...prev, { id: nanoid(), fromNodeId: node.id, toNodeId: childId }]);
setSelectedNodeIds(new Set([childId]));
setDialogNodeId(childId);
const controller = startGenerationRequest(childId, node.id, childId);
try {
const image = await requestEdit(generationConfig, prompt, [{ id: node.id, name: `${node.title || node.id}.png`, type: node.metadata.mimeType || "image/png", dataUrl: node.metadata.content, storageKey: node.metadata.storageKey }]).then(
const image = await requestEdit(generationConfig, prompt, [{ id: node.id, name: `${node.title || node.id}.png`, type: node.metadata.mimeType || "image/png", dataUrl: node.metadata.content, storageKey: node.metadata.storageKey }], undefined, { signal: controller.signal }).then(
(items) => items[0],
);
const uploaded = await uploadImage(image.dataUrl);
const size = fitNodeSize(uploaded.width, uploaded.height, imageConfig.width, imageConfig.height);
setNodes((prev) => prev.map((item) => (item.id === childId ? { ...item, width: size.width, height: size.height, metadata: { ...item.metadata, ...imageMetadata(uploaded), prompt, ...generationMetadata } } : item)));
} catch (error) {
if (isGenerationCanceled(error)) return;
const errorDetails = error instanceof Error ? error.message : "生成失败";
setNodes((prev) => prev.map((item) => (item.id === childId ? { ...item, metadata: { ...item.metadata, status: NODE_STATUS_ERROR, errorDetails } } : item)));
} finally {
finishGenerationRequest(childId, controller);
setRunningNodeId(null);
}
},
[effectiveConfig, openConfigDialog],
[effectiveConfig, finishGenerationRequest, openConfigDialog, startGenerationRequest],
);
const handleFontSizeChange = useCallback((nodeId: string, fontSize: number) => {
@@ -1880,15 +1941,22 @@ function InfiniteCanvasPage() {
}
setRunningNodeId(nodeId);
const runController = startGenerationRequest(nodeId, nodeId, nodeId);
const sourceTextContent = sourceNode?.type === CanvasNodeType.Text ? sourceNode.metadata?.content?.trim() || "" : "";
const editingTextNode = mode === "text" && Boolean(sourceTextContent);
const generationContext = await hydrateNodeGenerationContext(
buildNodeGenerationContext(nodeId, nodesRef.current, connectionsRef.current, editingTextNode ? `请根据要求修改以下文本。\n\n原文:\n${sourceTextContent}\n\n修改要求:\n${prompt}` : prompt),
);
const effectivePrompt = generationContext.prompt.trim();
if (runController.signal.aborted) {
finishGenerationRequest(nodeId, runController);
setRunningNodeId(null);
return;
}
const markSourceStatus = sourceNode?.type !== CanvasNodeType.Image && !editingTextNode;
const statusPrompt = sourceNode?.type === CanvasNodeType.Config ? effectivePrompt : prompt;
if (!effectivePrompt && (mode === "text" || mode === "audio")) {
finishGenerationRequest(nodeId, runController);
setRunningNodeId(null);
return;
}
@@ -1991,14 +2059,17 @@ function InfiniteCanvasPage() {
setSelectedConnectionId(null);
setDialogNodeId(nodeId);
const controller = runController;
targetIds.forEach((targetId) => startGenerationRequest(targetId, nodeId, nodeId, controller));
if (count > 1) startGenerationRequest(rootId, nodeId, nodeId, controller);
let hasSuccess = false;
let hasFailure = false;
await Promise.all(
targetIds.map(async (targetId) => {
try {
const image = referenceImages.length
? await requestEdit({ ...generationConfig, count: "1" }, effectivePrompt, referenceImages).then((items) => items[0])
: await requestGeneration({ ...generationConfig, count: "1" }, effectivePrompt).then((items) => items[0]);
? await requestEdit({ ...generationConfig, count: "1" }, effectivePrompt, referenceImages, undefined, { signal: controller.signal }).then((items) => items[0])
: await requestGeneration({ ...generationConfig, count: "1" }, effectivePrompt, { signal: controller.signal }).then((items) => items[0]);
const uploaded = await uploadImage(image.dataUrl);
const imageSize = fitNodeSize(uploaded.width, uploaded.height, imageConfig.width, imageConfig.height);
setNodes((prev) => {
@@ -2029,13 +2100,21 @@ function InfiniteCanvasPage() {
if (isConfigNode) setNodes((prev) => prev.map((node) => (node.id === nodeId ? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_SUCCESS, errorDetails: undefined } } : node)));
return true;
} catch (error) {
if (isGenerationCanceled(error)) return false;
const errorDetails = error instanceof Error ? error.message : "生成失败";
hasFailure = true;
setNodes((prev) => prev.map((node) => (node.id === targetId ? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_ERROR, errorDetails } } : node)));
return false;
} finally {
finishGenerationRequest(targetId, controller);
}
return false;
}),
);
if (count > 1) finishGenerationRequest(rootId, controller);
if (controller.signal.aborted) {
setNodes((prev) => prev.map((node) => (node.id === nodeId && isConfigNode && node.metadata?.status === NODE_STATUS_LOADING ? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_IDLE, errorDetails: undefined } } : node)));
return;
}
if (hasFailure) message.error(hasSuccess ? "部分图片生成失败" : "全部图片生成失败");
setNodes((prev) =>
prev.map((node) =>
@@ -2068,9 +2147,14 @@ function InfiniteCanvasPage() {
pendingChildIds = [videoId];
setNodes((prev) => (isEmptyVideoNode ? prev.map((node) => (node.id === nodeId ? { ...node, ...videoNode } : node)) : [...prev.map((node) => (node.id === nodeId ? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_SUCCESS } } : node)), videoNode]));
if (!isEmptyVideoNode) setConnections((prev) => [...prev, { id: nanoid(), fromNodeId: nodeId, toNodeId: videoId }]);
const video = await storeGeneratedVideo(await requestVideoGeneration(generationConfig, effectivePrompt, generationContext.referenceImages, generationContext.referenceVideos, generationContext.referenceAudios));
const videoSize = fitNodeSize(video.width || spec.width, video.height || spec.height, VIDEO_NODE_MAX_WIDTH, VIDEO_NODE_MAX_HEIGHT);
setNodes((prev) => prev.map((node) => (node.id === videoId ? { ...node, width: videoSize.width, height: videoSize.height, position: { x: node.position.x + node.width / 2 - videoSize.width / 2, y: node.position.y + node.height / 2 - videoSize.height / 2 }, metadata: { ...node.metadata, ...videoMetadata(video), prompt: effectivePrompt, model: generationConfig.model, size: generationConfig.size, seconds: generationConfig.videoSeconds, vquality: generationConfig.vquality, generateAudio: generationConfig.videoGenerateAudio, watermark: generationConfig.videoWatermark, references: generationReferenceUrls(generationContext) } } : node)));
const controller = startGenerationRequest(videoId, nodeId, nodeId, runController);
try {
const video = await storeGeneratedVideo(await requestVideoGeneration(generationConfig, effectivePrompt, generationContext.referenceImages, generationContext.referenceVideos, generationContext.referenceAudios, { signal: controller.signal }));
const videoSize = fitNodeSize(video.width || spec.width, video.height || spec.height, VIDEO_NODE_MAX_WIDTH, VIDEO_NODE_MAX_HEIGHT);
setNodes((prev) => prev.map((node) => (node.id === videoId ? { ...node, width: videoSize.width, height: videoSize.height, position: { x: node.position.x + node.width / 2 - videoSize.width / 2, y: node.position.y + node.height / 2 - videoSize.height / 2 }, metadata: { ...node.metadata, ...videoMetadata(video), prompt: effectivePrompt, model: generationConfig.model, size: generationConfig.size, seconds: generationConfig.videoSeconds, vquality: generationConfig.vquality, generateAudio: generationConfig.videoGenerateAudio, watermark: generationConfig.videoWatermark, references: generationReferenceUrls(generationContext) } } : node)));
} finally {
finishGenerationRequest(videoId, controller);
}
return;
}
@@ -2091,8 +2175,13 @@ function InfiniteCanvasPage() {
pendingChildIds = [audioId];
setNodes((prev) => (isEmptyAudioNode ? prev.map((node) => (node.id === nodeId ? { ...node, ...audioNode } : node)) : [...prev.map((node) => (node.id === nodeId ? { ...node, metadata: { ...node.metadata, status: NODE_STATUS_SUCCESS } } : node)), audioNode]));
if (!isEmptyAudioNode) setConnections((prev) => [...prev, { id: nanoid(), fromNodeId: nodeId, toNodeId: audioId }]);
const audio = await storeGeneratedAudio(await requestAudioGeneration(generationConfig, effectivePrompt), generationConfig.audioFormat);
setNodes((prev) => prev.map((node) => (node.id === audioId ? { ...node, metadata: { ...node.metadata, ...audioMetadata(audio), prompt: effectivePrompt, ...buildAudioGenerationMetadata(generationConfig) } } : node)));
const controller = startGenerationRequest(audioId, nodeId, nodeId, runController);
try {
const audio = await storeGeneratedAudio(await requestAudioGeneration(generationConfig, effectivePrompt, { signal: controller.signal }), generationConfig.audioFormat);
setNodes((prev) => prev.map((node) => (node.id === audioId ? { ...node, metadata: { ...node.metadata, ...audioMetadata(audio), prompt: effectivePrompt, ...buildAudioGenerationMetadata(generationConfig) } } : node)));
} finally {
finishGenerationRequest(audioId, controller);
}
return;
}
@@ -2121,17 +2210,21 @@ function InfiniteCanvasPage() {
setConnections((prev) => [...prev, ...childIds.map((childId) => ({ id: nanoid(), fromNodeId: nodeId, toNodeId: childId }))]);
}
const controller = runController;
const textTargetIds = childIds.length ? childIds : [nodeId];
textTargetIds.forEach((targetNodeId) => startGenerationRequest(targetNodeId, nodeId, nodeId, controller));
const answers = await Promise.all(
(childIds.length ? childIds : [nodeId]).map((targetNodeId) => {
textTargetIds.map((targetNodeId) => {
let localStreamed = "";
return requestImageQuestion(generationConfig, buildNodeResponseMessages({ ...generationContext, prompt: effectivePrompt }), (text) => {
localStreamed = text;
streamed = text;
if (isConfigNode) return;
setNodes((prev) => prev.map((node) => (node.id === targetNodeId ? { ...node, type: CanvasNodeType.Text, metadata: { ...node.metadata, content: text, status: NODE_STATUS_LOADING } } : node)));
}).then((answer) => ({ nodeId: targetNodeId, content: answer || localStreamed }));
}, { signal: controller.signal }).then((answer) => ({ nodeId: targetNodeId, content: answer || localStreamed })).finally(() => finishGenerationRequest(targetNodeId, controller));
}),
);
if (controller.signal.aborted) return;
const answerByNodeId = new Map(answers.map((item) => [item.nodeId, item.content]));
setNodes((prev) =>
prev.map((node) =>
@@ -2145,16 +2238,18 @@ function InfiniteCanvasPage() {
),
);
} catch (error) {
if (isGenerationCanceled(error)) return;
const errorDetails = error instanceof Error ? error.message : "生成失败";
message.error(errorDetails);
setNodes((prev) =>
prev.map((node) => (node.id === nodeId || pendingChildIds.includes(node.id) ? (node.id === nodeId && !markSourceStatus ? node : { ...node, metadata: { ...node.metadata, status: NODE_STATUS_ERROR, errorDetails } }) : node)),
);
} finally {
finishGenerationRequest(nodeId, runController);
setRunningNodeId(null);
}
},
[effectiveConfig, openConfigDialog],
[effectiveConfig, finishGenerationRequest, isAiConfigReady, message, openConfigDialog, startGenerationRequest],
);
useEffect(() => {
generateNodeRef.current = handleGenerateNode;
@@ -2200,6 +2295,7 @@ function InfiniteCanvasPage() {
setRunningNodeId(node.id);
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, metadata: { ...item.metadata, status: NODE_STATUS_LOADING, errorDetails: undefined } } : item)));
const controller = startGenerationRequest(node.id, sourceNode.id, node.id);
try {
if (node.type === CanvasNodeType.Text) {
@@ -2208,23 +2304,23 @@ function InfiniteCanvasPage() {
const answer = await requestImageQuestion(generationConfig, buildNodeResponseMessages({ ...context, prompt }), (text) => {
streamed = text;
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, type: CanvasNodeType.Text, metadata: { ...item.metadata, content: text, status: NODE_STATUS_LOADING } } : item)));
});
}, { signal: controller.signal });
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, type: CanvasNodeType.Text, metadata: { ...item.metadata, content: answer || streamed, prompt, status: NODE_STATUS_SUCCESS } } : item)));
return;
}
if (node.type === CanvasNodeType.Video) {
const video = await storeGeneratedVideo(await requestVideoGeneration(generationConfig, prompt, retryImages, context?.referenceVideos || [], context?.referenceAudios || []));
const video = await storeGeneratedVideo(await requestVideoGeneration(generationConfig, prompt, retryImages, context?.referenceVideos || [], context?.referenceAudios || [], { signal: controller.signal }));
const videoSize = fitNodeSize(video.width || node.width, video.height || node.height, VIDEO_NODE_MAX_WIDTH, VIDEO_NODE_MAX_HEIGHT);
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, width: videoSize.width, height: videoSize.height, position: { x: item.position.x + item.width / 2 - videoSize.width / 2, y: item.position.y + item.height / 2 - videoSize.height / 2 }, metadata: { ...item.metadata, ...videoMetadata(video), prompt, model: generationConfig.model, size: generationConfig.size, seconds: generationConfig.videoSeconds, vquality: generationConfig.vquality, generateAudio: generationConfig.videoGenerateAudio, watermark: generationConfig.videoWatermark } } : item)));
return;
}
if (node.type === CanvasNodeType.Audio) {
const audio = await storeGeneratedAudio(await requestAudioGeneration(generationConfig, prompt), generationConfig.audioFormat);
const audio = await storeGeneratedAudio(await requestAudioGeneration(generationConfig, prompt, { signal: controller.signal }), generationConfig.audioFormat);
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, metadata: { ...item.metadata, ...audioMetadata(audio), prompt, ...buildAudioGenerationMetadata(generationConfig) } } : item)));
return;
}
const image = useReferenceImages ? await requestEdit(generationConfig, prompt, retryImages).then((items) => items[0]) : await requestGeneration(generationConfig, prompt).then((items) => items[0]);
const image = useReferenceImages ? await requestEdit(generationConfig, prompt, retryImages, undefined, { signal: controller.signal }).then((items) => items[0]) : await requestGeneration(generationConfig, prompt, { signal: controller.signal }).then((items) => items[0]);
const uploadedImage = await uploadImage(image.dataUrl);
const imageConfig = NODE_DEFAULT_SIZE[CanvasNodeType.Image];
const imageSize = fitNodeSize(uploadedImage.width, uploadedImage.height, imageConfig.width, imageConfig.height);
@@ -2245,14 +2341,16 @@ function InfiniteCanvasPage() {
),
);
} catch (error) {
if (isGenerationCanceled(error)) return;
const errorDetails = error instanceof Error ? error.message : "生成失败";
message.error(errorDetails);
setNodes((prev) => prev.map((item) => (item.id === node.id ? { ...item, metadata: { ...item.metadata, status: NODE_STATUS_ERROR, errorDetails } } : item)));
} finally {
finishGenerationRequest(node.id, controller);
setRunningNodeId(null);
}
},
[effectiveConfig, message, openConfigDialog],
[effectiveConfig, finishGenerationRequest, isAiConfigReady, message, openConfigDialog, startGenerationRequest],
);
const generateImageFromTextNode = useCallback(
@@ -2484,6 +2582,7 @@ function InfiniteCanvasPage() {
onPromptChange={handleNodePromptChange}
onConfigChange={handleConfigNodeChange}
onGenerate={handleGenerateNode}
onStop={confirmStopGeneration}
onImageSettingsOpenChange={(open) => {
setNodeImageSettingsOpen(open);
if (open) setToolbarNodeId(null);
@@ -2498,6 +2597,7 @@ function InfiniteCanvasPage() {
inputSummary={getInputSummary(configInputsById.get(contentNode.id) || [])}
onConfigChange={handleConfigNodeChange}
onComposerToggle={() => setDialogNodeId((current) => (current === contentNode.id ? null : contentNode.id))}
onStop={confirmStopGeneration}
onGenerate={(nodeId) => {
const target = nodesRef.current.find((item) => item.id === nodeId);
void handleGenerateNode(nodeId, target?.metadata?.generationMode || "image", target?.metadata?.composerContent ?? target?.metadata?.prompt ?? "");
@@ -3041,6 +3141,10 @@ function resetInterruptedGeneration(nodes: CanvasNodeData[]) {
return nodes.map((node) => (node.metadata?.status === "loading" ? { ...node, metadata: { ...node.metadata, status: "error" as const, errorDetails: "页面刷新后生成已中断,请重新生成。" } } : node));
}
function isGenerationCanceled(error: unknown) {
return error instanceof Error && (error.message === "请求已取消" || error.name === "AbortError");
}
function findRetrySourceNode(nodeId: string, nodes: CanvasNodeData[], connections: CanvasConnection[]) {
const queue = connections.filter((connection) => connection.toNodeId === nodeId).map((connection) => connection.fromNodeId);
const visited = new Set<string>();
@@ -1,7 +1,7 @@
"use client";
import type { CSSProperties } from "react";
import { Image as ImageIcon, LoaderCircle, MessageSquare, Music2, Play, Settings2, Video } from "lucide-react";
import { Image as ImageIcon, LoaderCircle, MessageSquare, Music2, Play, Settings2, Square, Video } from "lucide-react";
import { Button, Segmented } from "antd";
import { ModelPicker } from "@/components/model-picker";
@@ -20,10 +20,11 @@ type CanvasConfigNodePanelProps = {
inputSummary: { textCount: number; imageCount: number; videoCount: number; audioCount: number };
onConfigChange: (nodeId: string, patch: Partial<CanvasNodeMetadata>) => void;
onGenerate: (nodeId: string) => void;
onStop: (nodeId: string) => void;
onComposerToggle: () => void;
};
export function CanvasConfigNodePanel({ node, isRunning, inputSummary, onConfigChange, onGenerate, onComposerToggle }: CanvasConfigNodePanelProps) {
export function CanvasConfigNodePanel({ node, isRunning, inputSummary, onConfigChange, onGenerate, onStop, onComposerToggle }: CanvasConfigNodePanelProps) {
const globalConfig = useEffectiveConfig();
const openConfigDialog = useConfigStore((state) => state.openConfigDialog);
const theme = canvasThemes[useThemeStore((state) => state.theme)];
@@ -113,17 +114,28 @@ export function CanvasConfigNodePanel({ node, isRunning, inputSummary, onConfigC
<Button
type="primary"
className="mt-auto !h-9 !w-full !cursor-pointer !rounded-lg"
disabled={isRunning || !canGenerate}
danger={isRunning}
disabled={!isRunning && !canGenerate}
onMouseDown={(event) => event.stopPropagation()}
onClick={() => onGenerate(node.id)}
onClick={() => (isRunning ? onStop(node.id) : onGenerate(node.id))}
>
<span className="inline-flex items-center gap-1.5">
<span className="inline-flex items-center gap-1">
<CreditSymbol />
{credits.toLocaleString()}
</span>
{isRunning ? <LoaderCircle className="size-4 animate-spin" /> : <Play className="size-4" />}
<span></span>
{isRunning ? (
<>
<LoaderCircle className="size-4 animate-spin" />
<Square className="size-3.5 fill-current" />
<span></span>
</>
) : (
<>
<span className="inline-flex items-center gap-1">
<CreditSymbol />
{credits.toLocaleString()}
</span>
<Play className="size-4" />
<span></span>
</>
)}
</span>
</Button>
</div>
@@ -1,7 +1,7 @@
"use client";
import { useEffect, useState } from "react";
import { ArrowUp, LoaderCircle } from "lucide-react";
import { ArrowUp, LoaderCircle, Square } from "lucide-react";
import { Button } from "antd";
import { ModelPicker } from "@/components/model-picker";
@@ -25,11 +25,12 @@ type CanvasNodePromptPanelProps = {
onPromptChange: (nodeId: string, prompt: string) => void;
onConfigChange: (nodeId: string, patch: Partial<CanvasNodeData["metadata"]>) => void;
onGenerate: (nodeId: string, mode: CanvasNodeGenerationMode, prompt: string) => void;
onStop: (nodeId: string) => void;
mentionReferences?: CanvasResourceReference[];
onImageSettingsOpenChange?: (open: boolean) => void;
};
export function CanvasNodePromptPanel({ node, isRunning, onPromptChange, onConfigChange, onGenerate, mentionReferences = [], onImageSettingsOpenChange }: CanvasNodePromptPanelProps) {
export function CanvasNodePromptPanel({ node, isRunning, onPromptChange, onConfigChange, onGenerate, onStop, mentionReferences = [], onImageSettingsOpenChange }: CanvasNodePromptPanelProps) {
const globalConfig = useEffectiveConfig();
const openConfigDialog = useConfigStore((state) => state.openConfigDialog);
const theme = canvasThemes[useThemeStore((state) => state.theme)];
@@ -107,16 +108,27 @@ export function CanvasNodePromptPanel({ node, isRunning, onPromptChange, onConfi
<Button
type="primary"
className="!h-10 !min-w-16 shrink-0 !rounded-full !px-3"
disabled={isRunning || !prompt.trim()}
onClick={submit}
aria-label="生成"
danger={isRunning}
disabled={!isRunning && !prompt.trim()}
onClick={() => (isRunning ? onStop(node.id) : submit())}
aria-label={isRunning ? "停止生成" : "生成"}
>
<span className="flex items-center gap-1.5">
<span className="inline-flex items-center gap-1 text-xs font-medium tabular-nums">
<CreditSymbol />
{credits.toLocaleString()}
</span>
{isRunning ? <LoaderCircle className="size-4 animate-spin" /> : <ArrowUp className="size-4" />}
{isRunning ? (
<>
<LoaderCircle className="size-4 animate-spin" />
<Square className="size-3.5 fill-current" />
<span className="text-xs font-medium"></span>
</>
) : (
<>
<span className="inline-flex items-center gap-1 text-xs font-medium tabular-nums">
<CreditSymbol />
{credits.toLocaleString()}
</span>
<ArrowUp className="size-4" />
</>
)}
</span>
</Button>
</div>
+5 -2
View File
@@ -4,6 +4,8 @@ import { audioMimeType, normalizeAudioFormatValue, normalizeAudioSpeedValue, nor
import { uploadMediaFile, type UploadedFile } from "@/services/file-storage";
import { buildApiUrl, resolveModelRequestConfig, type AiConfig } from "@/stores/use-config-store";
type RequestOptions = { signal?: AbortSignal };
function aiApiUrl(config: AiConfig, path: string) {
return buildApiUrl(config.baseUrl, path);
}
@@ -15,7 +17,7 @@ function aiHeaders(config: AiConfig) {
};
}
export async function requestAudioGeneration(config: AiConfig, prompt: string): Promise<Blob> {
export async function requestAudioGeneration(config: AiConfig, prompt: string, options?: RequestOptions): Promise<Blob> {
const requestConfig = resolveModelRequestConfig(config, config.model || config.audioModel);
const model = requestConfig.model.trim();
assertAudioConfig(requestConfig, model);
@@ -33,7 +35,7 @@ export async function requestAudioGeneration(config: AiConfig, prompt: string):
speed: Number(normalizeAudioSpeedValue(config.audioSpeed)),
...(instructions ? { instructions } : {}),
},
{ headers: aiHeaders(requestConfig), responseType: "blob" },
{ headers: aiHeaders(requestConfig), responseType: "blob", signal: options?.signal },
);
await assertAudioBlob(response.data);
return response.data.type.startsWith("audio/") ? response.data : new Blob([response.data], { type: audioMimeType(format) });
@@ -66,6 +68,7 @@ async function assertAudioBlob(blob: Blob) {
}
function readAxiosError(error: unknown, fallback: string) {
if (axios.isCancel(error)) return "请求已取消";
if (axios.isAxiosError<{ error?: { message?: string }; msg?: string; code?: number }>(error)) {
const responseData = error.response?.data;
return responseData?.msg || responseData?.error?.message || statusMessage(error.response?.status, fallback);
+13 -8
View File
@@ -71,6 +71,7 @@ type ImageApiResponse = {
code?: number;
msg?: string;
};
type RequestOptions = { signal?: AbortSignal };
const QUALITY_BASE: Record<string, number> = {
low: 1024,
@@ -188,10 +189,12 @@ function parseImagePayload(payload: ImageApiResponse) {
}
function readAxiosError(error: unknown, fallback: string) {
if (axios.isCancel(error)) return "请求已取消";
if (axios.isAxiosError<{ error?: { message?: string }; msg?: string; code?: number }>(error)) {
const responseData = error.response?.data;
return responseData?.msg || responseData?.error?.message || readStatusError(error.response?.status, fallback);
}
if (error instanceof DOMException && error.name === "AbortError") return "请求已取消";
return error instanceof Error ? error.message : fallback;
}
@@ -336,11 +339,12 @@ function consumeResponseStreamText(state: ResponseStreamState, text: string, onD
}
}
async function requestStreamingResponse(config: AiConfig, body: Record<string, unknown>, onDelta?: (text: string) => void): Promise<ToolResponseResult> {
async function requestStreamingResponse(config: AiConfig, body: Record<string, unknown>, onDelta?: (text: string) => void, options?: RequestOptions): Promise<ToolResponseResult> {
const response = await fetch(aiApiUrl(config, "/responses"), {
method: "POST",
headers: { ...aiHeaders(config, "application/json"), Accept: "text/event-stream" },
body: JSON.stringify({ ...body, stream: true }),
signal: options?.signal,
});
if (!response.ok) throw new Error(await readFetchError(response, "请求失败"));
if (!response.body) {
@@ -366,7 +370,7 @@ async function requestStreamingResponse(config: AiConfig, body: Record<string, u
return { ...result, content: state.text || result.content };
}
export async function requestGeneration(config: AiConfig, prompt: string) {
export async function requestGeneration(config: AiConfig, prompt: string, options?: RequestOptions) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.imageModel);
const n = Math.max(1, Math.min(15, Math.floor(Math.abs(Number(config.count)) || 1)));
const quality = normalizeQuality(config.quality);
@@ -385,6 +389,7 @@ export async function requestGeneration(config: AiConfig, prompt: string) {
},
{
headers: aiHeaders(requestConfig, "application/json"),
signal: options?.signal,
},
);
const images = parseImagePayload(response.data);
@@ -394,7 +399,7 @@ export async function requestGeneration(config: AiConfig, prompt: string) {
}
}
export async function requestEdit(config: AiConfig, prompt: string, references: ReferenceImage[], mask?: ReferenceImage) {
export async function requestEdit(config: AiConfig, prompt: string, references: ReferenceImage[], mask?: ReferenceImage, options?: RequestOptions) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.imageModel);
const n = Math.max(1, Math.min(15, Math.floor(Math.abs(Number(config.count)) || 1)));
const quality = normalizeQuality(config.quality);
@@ -417,7 +422,7 @@ export async function requestEdit(config: AiConfig, prompt: string, references:
if (mask) formData.set("mask", dataUrlToFile(mask));
try {
const response = await axios.post<ImageApiResponse>(aiApiUrl(requestConfig, "/images/edits"), formData, { headers: aiHeaders(requestConfig) });
const response = await axios.post<ImageApiResponse>(aiApiUrl(requestConfig, "/images/edits"), formData, { headers: aiHeaders(requestConfig), signal: options?.signal });
const images = parseImagePayload(response.data);
return images;
} catch (error) {
@@ -425,13 +430,13 @@ export async function requestEdit(config: AiConfig, prompt: string, references:
}
}
export async function requestImageQuestion(config: AiConfig, messages: AiTextMessage[], onDelta: (text: string) => void) {
export async function requestImageQuestion(config: AiConfig, messages: AiTextMessage[], onDelta: (text: string) => void, options?: RequestOptions) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.textModel);
try {
const answer = (await requestStreamingResponse(requestConfig, {
model: requestConfig.model,
input: toResponseInput(withSystemMessage(requestConfig, messages)),
}, onDelta)).content || "没有返回内容";
}, onDelta, options)).content || "没有返回内容";
if (answer === "没有返回内容") onDelta(answer);
return answer;
} catch (error) {
@@ -439,7 +444,7 @@ export async function requestImageQuestion(config: AiConfig, messages: AiTextMes
}
}
export async function requestToolResponse(config: AiConfig, messages: ResponseInputMessage[], tools: ResponseFunctionTool[], toolChoice: ToolChoice = "auto", onDelta?: (text: string) => void): Promise<ToolResponseResult> {
export async function requestToolResponse(config: AiConfig, messages: ResponseInputMessage[], tools: ResponseFunctionTool[], toolChoice: ToolChoice = "auto", onDelta?: (text: string) => void, options?: RequestOptions): Promise<ToolResponseResult> {
const requestConfig = resolveModelRequestConfig(config, config.model || config.textModel);
try {
return await requestStreamingResponse(requestConfig, {
@@ -448,7 +453,7 @@ export async function requestToolResponse(config: AiConfig, messages: ResponseIn
tools: tools.map(toResponseTool),
tool_choice: toolChoice,
parallel_tool_calls: false,
}, onDelta);
}, onDelta, options);
} catch (error) {
throw new Error(readAxiosError(error, "请求失败"));
}
+43 -24
View File
@@ -17,6 +17,7 @@ type SeedanceTask = {
content?: { video_url?: string; last_frame_url?: string } | null;
};
type ApiEnvelope<T> = T | { code?: number; data?: T | null; msg?: string };
type RequestOptions = { signal?: AbortSignal };
export type VideoGenerationResult = { blob?: Blob; url?: string; mimeType?: string };
export type VideoGenerationTask = { id: string; provider: "openai" | "seedance"; model: string };
@@ -33,36 +34,37 @@ function aiHeaders(config: AiConfig, contentType?: string) {
};
}
export async function requestVideoGeneration(config: AiConfig, prompt: string, references: ReferenceImage[] = [], videoReferences: ReferenceVideo[] = [], audioReferences: ReferenceAudio[] = []): Promise<VideoGenerationResult> {
const task = await createVideoGenerationTask(config, prompt, references, videoReferences, audioReferences);
export async function requestVideoGeneration(config: AiConfig, prompt: string, references: ReferenceImage[] = [], videoReferences: ReferenceVideo[] = [], audioReferences: ReferenceAudio[] = [], options?: RequestOptions): Promise<VideoGenerationResult> {
const task = await createVideoGenerationTask(config, prompt, references, videoReferences, audioReferences, options);
const delayMs = task.provider === "seedance" ? 5000 : 2500;
for (let attempt = 0; attempt < 120; attempt += 1) {
const state = await pollVideoGenerationTask(config, task);
if (options?.signal?.aborted) throw new DOMException("Aborted", "AbortError");
const state = await pollVideoGenerationTask(config, task, options);
if (state.status === "completed") return state.result;
if (state.status === "failed") throw new Error(state.error);
if (attempt === 119) throw new Error(`${task.provider === "seedance" ? "Seedance " : ""}视频生成超时,请稍后重试`);
await delay(delayMs);
await delay(delayMs, options?.signal);
}
throw new Error("视频生成超时,请稍后重试");
}
export async function createVideoGenerationTask(config: AiConfig, prompt: string, references: ReferenceImage[] = [], videoReferences: ReferenceVideo[] = [], audioReferences: ReferenceAudio[] = []): Promise<VideoGenerationTask> {
export async function createVideoGenerationTask(config: AiConfig, prompt: string, references: ReferenceImage[] = [], videoReferences: ReferenceVideo[] = [], audioReferences: ReferenceAudio[] = [], options?: RequestOptions): Promise<VideoGenerationTask> {
const selectedModel = (config.model || config.videoModel).trim();
const requestConfig = resolveModelRequestConfig(config, selectedModel);
assertVideoConfig(requestConfig, requestConfig.model);
if (isSeedanceVideoConfig(requestConfig)) {
return createSeedanceTask(requestConfig, selectedModel, prompt, references, videoReferences, audioReferences);
return createSeedanceTask(requestConfig, selectedModel, prompt, references, videoReferences, audioReferences, options);
}
if (videoReferences.length || audioReferences.length) {
throw new Error("当前视频接口不支持参考视频或参考音频,请切换到 Seedance 2.0 / 火山 Agent Plan 模型,或移除参考素材");
}
return createOpenAIVideoTask(requestConfig, selectedModel, prompt, references);
return createOpenAIVideoTask(requestConfig, selectedModel, prompt, references, options);
}
export async function pollVideoGenerationTask(config: AiConfig, task: VideoGenerationTask): Promise<VideoGenerationTaskState> {
export async function pollVideoGenerationTask(config: AiConfig, task: VideoGenerationTask, options?: RequestOptions): Promise<VideoGenerationTaskState> {
const requestConfig = resolveModelRequestConfig(config, task.model);
assertVideoConfig(requestConfig, requestConfig.model);
return task.provider === "seedance" ? pollSeedanceTask(requestConfig, task) : pollOpenAIVideoTask(requestConfig, task);
return task.provider === "seedance" ? pollSeedanceTask(requestConfig, task, options) : pollOpenAIVideoTask(requestConfig, task, options);
}
export async function storeGeneratedVideo(result: VideoGenerationResult): Promise<UploadedFile> {
@@ -71,7 +73,7 @@ export async function storeGeneratedVideo(result: VideoGenerationResult): Promis
throw new Error("视频接口没有返回可播放的视频");
}
async function createOpenAIVideoTask(config: AiConfig, model: string, prompt: string, references: ReferenceImage[]): Promise<VideoGenerationTask> {
async function createOpenAIVideoTask(config: AiConfig, model: string, prompt: string, references: ReferenceImage[], options?: RequestOptions): Promise<VideoGenerationTask> {
const body = new FormData();
body.append("model", modelOptionName(model));
body.append("prompt", prompt);
@@ -82,7 +84,7 @@ async function createOpenAIVideoTask(config: AiConfig, model: string, prompt: st
const files = await Promise.all(references.slice(0, 7).map(async (image) => dataUrlToFile({ ...image, dataUrl: await imageToDataUrl(image) })));
files.forEach((file) => body.append("input_reference[]", file));
try {
const created = unwrapVideoResponse((await axios.post<ApiVideoResponse>(aiApiUrl(config, "/videos"), body, { headers: aiHeaders(config) })).data);
const created = unwrapVideoResponse((await axios.post<ApiVideoResponse>(aiApiUrl(config, "/videos"), body, { headers: aiHeaders(config), signal: options?.signal })).data);
if (!created.id) throw new Error("视频接口没有返回任务 ID");
return { id: created.id, provider: "openai", model };
} catch (error) {
@@ -90,11 +92,11 @@ async function createOpenAIVideoTask(config: AiConfig, model: string, prompt: st
}
}
async function pollOpenAIVideoTask(config: AiConfig, task: VideoGenerationTask): Promise<VideoGenerationTaskState> {
async function pollOpenAIVideoTask(config: AiConfig, task: VideoGenerationTask, options?: RequestOptions): Promise<VideoGenerationTaskState> {
try {
const video = unwrapVideoResponse((await axios.get<ApiVideoResponse>(aiApiUrl(config, `/videos/${task.id}`), { headers: aiHeaders(config) })).data);
const video = unwrapVideoResponse((await axios.get<ApiVideoResponse>(aiApiUrl(config, `/videos/${task.id}`), { headers: aiHeaders(config), signal: options?.signal })).data);
if (video.status === "completed") {
const content = await axios.get<Blob>(aiApiUrl(config, `/videos/${task.id}/content`), { headers: aiHeaders(config), responseType: "blob" });
const content = await axios.get<Blob>(aiApiUrl(config, `/videos/${task.id}/content`), { headers: aiHeaders(config), responseType: "blob", signal: options?.signal });
await assertVideoBlob(content.data);
return { status: "completed", result: { blob: content.data } };
}
@@ -105,7 +107,7 @@ async function pollOpenAIVideoTask(config: AiConfig, task: VideoGenerationTask):
}
}
async function createSeedanceTask(config: AiConfig, model: string, prompt: string, references: ReferenceImage[], videoReferences: ReferenceVideo[], audioReferences: ReferenceAudio[]): Promise<VideoGenerationTask> {
async function createSeedanceTask(config: AiConfig, model: string, prompt: string, references: ReferenceImage[], videoReferences: ReferenceVideo[], audioReferences: ReferenceAudio[], options?: RequestOptions): Promise<VideoGenerationTask> {
if (audioReferences.length && !references.length && !videoReferences.length) {
throw new Error("Seedance 参考音频不能单独使用,请同时添加参考图或参考视频");
}
@@ -124,7 +126,7 @@ async function createSeedanceTask(config: AiConfig, model: string, prompt: strin
};
try {
const created = unwrapSeedanceTask((await axios.post<ApiEnvelope<SeedanceTask>>(seedanceApiUrl(config), payload, { headers: aiHeaders(config, "application/json") })).data);
const created = unwrapSeedanceTask((await axios.post<ApiEnvelope<SeedanceTask>>(seedanceApiUrl(config), payload, { headers: aiHeaders(config, "application/json"), signal: options?.signal })).data);
if (!created.id) throw new Error("Seedance 接口没有返回任务 ID");
return { id: created.id, provider: "seedance", model };
} catch (error) {
@@ -132,13 +134,13 @@ async function createSeedanceTask(config: AiConfig, model: string, prompt: strin
}
}
async function pollSeedanceTask(config: AiConfig, task: VideoGenerationTask): Promise<VideoGenerationTaskState> {
async function pollSeedanceTask(config: AiConfig, task: VideoGenerationTask, options?: RequestOptions): Promise<VideoGenerationTaskState> {
try {
const state = unwrapSeedanceTask((await axios.get<ApiEnvelope<SeedanceTask>>(seedanceApiUrl(config, task.id), { headers: aiHeaders(config) })).data);
const state = unwrapSeedanceTask((await axios.get<ApiEnvelope<SeedanceTask>>(seedanceApiUrl(config, task.id), { headers: aiHeaders(config), signal: options?.signal })).data);
if (state.status === "succeeded") {
const url = state.content?.video_url;
if (!url) return { status: "failed", error: "Seedance 任务成功但没有返回视频 URL" };
return { status: "completed", result: await videoResultFromUrl(url) };
return { status: "completed", result: await videoResultFromUrl(url, options) };
}
if (state.status === "failed" || state.status === "cancelled" || state.status === "expired") return { status: "failed", error: state.error?.message || `Seedance 视频生成${state.status === "expired" ? "超时" : "失败"}` };
return { status: "pending" };
@@ -215,12 +217,13 @@ async function resolveSeedanceAudioUrl(audio: ReferenceAudio) {
return blobToDataUrl(blob);
}
async function videoResultFromUrl(url: string): Promise<VideoGenerationResult> {
async function videoResultFromUrl(url: string, options?: RequestOptions): Promise<VideoGenerationResult> {
try {
const response = await axios.get<Blob>(url, { responseType: "blob" });
const response = await axios.get<Blob>(url, { responseType: "blob", signal: options?.signal });
await assertVideoBlob(response.data);
return { blob: response.data };
} catch {
} catch (error) {
if (axios.isCancel(error) || options?.signal?.aborted) throw error;
return { url, mimeType: "video/mp4" };
}
}
@@ -269,10 +272,12 @@ function unwrapEnvelope<T>(payload: ApiEnvelope<T>, emptyMessage: string): T {
}
function readAxiosError(error: unknown, fallback: string) {
if (axios.isCancel(error)) return "请求已取消";
if (axios.isAxiosError<{ error?: { message?: string }; msg?: string; code?: number }>(error)) {
const responseData = error.response?.data;
return responseData?.msg || responseData?.error?.message || statusMessage(error.response?.status, fallback);
}
if (error instanceof DOMException && error.name === "AbortError") return "请求已取消";
return error instanceof Error ? error.message : fallback;
}
@@ -298,8 +303,22 @@ function isPublicMediaUrl(value: string) {
return /^https?:\/\//i.test(value || "");
}
function delay(ms: number) {
return new Promise((resolve) => setTimeout(resolve, ms));
function delay(ms: number, signal?: AbortSignal) {
return new Promise<void>((resolve, reject) => {
if (signal?.aborted) {
reject(new DOMException("Aborted", "AbortError"));
return;
}
const timer = setTimeout(resolve, ms);
signal?.addEventListener(
"abort",
() => {
clearTimeout(timer);
reject(new DOMException("Aborted", "AbortError"));
},
{ once: true },
);
});
}
function blobToDataUrl(blob: Blob) {