From 51a54dbdb26e9fedd5fd7d5cf8090b027ad6ac8b Mon Sep 17 00:00:00 2001 From: syntaxbullet Date: Fri, 10 Jul 2026 23:54:52 +0200 Subject: [PATCH] feat: implement generation workflow with resource loading and candidate management --- app/app.ts | 3 + operations/generation/options.test.ts | 23 ++++ operations/generation/options.ts | 41 ++++++ operations/generation/workflow.test.ts | 88 ++++++++++++ operations/generation/workflow.ts | 125 ++++++++++++++++++ view/App.tsx | 6 +- view/BottomControlsIsland.tsx | 6 +- .../GenerateActionControls.tsx | 85 ++---------- view/bottom-controls/GenerateControls.tsx | 48 ++----- 9 files changed, 308 insertions(+), 117 deletions(-) create mode 100644 operations/generation/options.test.ts create mode 100644 operations/generation/options.ts create mode 100644 operations/generation/workflow.test.ts create mode 100644 operations/generation/workflow.ts diff --git a/app/app.ts b/app/app.ts index 2d87064..40b5120 100644 --- a/app/app.ts +++ b/app/app.ts @@ -13,12 +13,14 @@ import { editorCommands } from "@commands/editor"; import { projectCommands } from "@commands/project"; import { createInitialAppState } from "@editor/initial-state"; import { createAppStore } from "@editor/store"; +import { createGenerationWorkflow } from "@operations/generation/workflow"; export type ImageStudioApp = ReturnType; export function createImageStudioApp(options?: { documentName?: string; createDefaultArtboard?: boolean }) { const registry = createCommandRegistry([...projectCommands, ...viewportCommands, ...selectionCommands, ...documentCommands, ...toolCommands, ...generationCommands, ...transformCommands, ...historyCommands, ...commandPaletteCommands, ...workspaceCommands, ...editorCommands]); const store = createAppStore(createInitialAppState(options?.documentName), registry); + const generation = createGenerationWorkflow(store); if (options?.createDefaultArtboard !== false) { const artboardId = crypto.randomUUID(); @@ -33,5 +35,6 @@ export function createImageStudioApp(options?: { documentName?: string; createDe return { registry, store, + workflows: { generation }, }; } diff --git a/operations/generation/options.test.ts b/operations/generation/options.test.ts new file mode 100644 index 0000000..53c0207 --- /dev/null +++ b/operations/generation/options.test.ts @@ -0,0 +1,23 @@ +import { describe, expect, test } from "bun:test"; +import { initialToolState } from "@editor/tools"; +import { resolveGenerationModeOptions, resolveGenerationModelOptions } from "./options"; + +describe("generation compatibility options", () => { + test("uses backend compatibility data for modes", () => { + const settings = { ...initialToolState.generate, architecture: "anima" as const, mode: "inpaint" as const }; + const modes = [ + { value: "text-to-image" as const, label: "Text" }, + { value: "inpaint" as const, label: "Inpaint" }, + ]; + const options = { architectures: [{ value: "anima" as const, label: "Anima", models: [], defaultModel: "auto", supportedModes: ["text-to-image" as const] }], models: [], samplers: [], schedulers: [], textEncoders: [], vaes: [] }; + + expect(resolveGenerationModeOptions(settings, options, modes).map((option) => option.value)).toEqual(["inpaint", "text-to-image"]); + }); + + test("keeps the current model alongside discovered models", () => { + const settings = { ...initialToolState.generate, model: "current.safetensors" }; + const options = { architectures: [], models: ["found.safetensors"], samplers: [], schedulers: [], textEncoders: [], vaes: [] }; + + expect(resolveGenerationModelOptions(settings, options).map((option) => option.value)).toEqual(["auto", "found.safetensors", "current.safetensors"]); + }); +}); diff --git a/operations/generation/options.ts b/operations/generation/options.ts new file mode 100644 index 0000000..8863718 --- /dev/null +++ b/operations/generation/options.ts @@ -0,0 +1,41 @@ +import type { GenerationOptions } from "@editor/state"; +import { generateArchitectureDefaults } from "@editor/tools"; +import type { GenerateMode, GenerateModel, GenerateSettings } from "@editor/tools"; + +export type GenerationSelectOption = { value: TValue; label: string }; + +export function resolveGenerationModelOptions(settings: GenerateSettings, options: GenerationOptions | undefined): readonly GenerationSelectOption[] { + const architecture = options?.architectures?.find((item) => item.value === settings.architecture); + const models = architecture?.models ?? (settings.architecture === "sdxl" ? options?.models : undefined) ?? []; + const fallbackModel = architecture?.defaultModel ?? generateArchitectureDefaults[settings.architecture].model; + const values = unique(["auto", ...models, ...(models.length === 0 && fallbackModel !== "auto" ? [fallbackModel] : []), settings.model]); + return values.map((model) => ({ value: model, label: model === "auto" ? "Auto" : model })); +} + +export function resolveGenerationSupportOptions(settings: GenerateSettings, options: GenerationOptions | undefined) { + const defaults = generateArchitectureDefaults[settings.architecture]; + return { + textEncoders: resolveGenerationStringOptions([...(options?.textEncoders ?? []), defaults.textEncoder].filter((value) => value !== "auto"), settings.textEncoder), + vaes: resolveGenerationStringOptions([...(options?.vaes ?? []), defaults.vae].filter((value) => value !== "auto"), settings.vae), + }; +} + +export function resolveGenerationStringOptions(values: string[] | undefined, current: string): readonly GenerationSelectOption[] { + return unique([...(values ?? []), current]).map((value) => ({ value, label: value })); +} + +export function resolveGenerationModeOptions( + settings: GenerateSettings, + options: GenerationOptions | undefined, + modes: readonly GenerationSelectOption[], +): readonly GenerationSelectOption[] { + const architecture = options?.architectures?.find((item) => item.value === settings.architecture); + const supportedModes = architecture?.supportedModes?.length ? architecture.supportedModes : generateArchitectureDefaults[settings.architecture].supportedModes; + const availableModes = modes.filter((mode) => supportedModes.includes(mode.value)); + if (availableModes.some((mode) => mode.value === settings.mode)) return availableModes; + return [modes.find((mode) => mode.value === settings.mode), ...availableModes].filter((mode): mode is GenerationSelectOption => Boolean(mode)); +} + +function unique(values: T[]): T[] { + return Array.from(new Set(values)); +} diff --git a/operations/generation/workflow.test.ts b/operations/generation/workflow.test.ts new file mode 100644 index 0000000..74ac004 --- /dev/null +++ b/operations/generation/workflow.test.ts @@ -0,0 +1,88 @@ +import { describe, expect, test } from "bun:test"; +import { documentCommands } from "@commands/document"; +import { generationCommands } from "@commands/generation"; +import { commandIds } from "@commands/ids"; +import { createCommandRegistry } from "@commands/registry"; +import { toolCommands } from "@commands/tool"; +import { createInitialAppState } from "@editor/initial-state"; +import { createAppStore } from "@editor/store"; +import { initialToolState } from "@editor/tools"; +import type { GenerationCandidate } from "@editor/state"; +import { createGenerationWorkflow, type GenerationWorkflowDependencies } from "./workflow"; + +describe("generation workflow", () => { + test("reads canonical application state when generation starts", async () => { + const app = createTestApp(); + let prompt = ""; + const workflow = createGenerationWorkflow(app.store, dependencies({ + runGenerate: async (options) => { prompt = options.settings.prompt; }, + })); + + app.store.dispatch(commandIds.toolSetGenerateSettings, { prompt: "Latest prompt" }); + await workflow.generate(); + + expect(prompt).toBe("Latest prompt"); + expect(app.store.getState().editor.generation.jobs[0]?.status).toBe("succeeded"); + }); + + test("uses one acceptance path to add a candidate as a layer", () => { + const app = createTestApp(); + const artboardId = app.store.getState().document.artboards[0]?.id; + if (!artboardId) throw new Error("Expected default artboard"); + app.store.dispatch(commandIds.generationAddCandidate, { candidate: candidate(artboardId) }); + const ids = ["asset-new", "layer-new"]; + const workflow = createGenerationWorkflow(app.store, dependencies({ createId: () => ids.shift() ?? "unused" })); + + workflow.applyCandidateAsLayer("candidate"); + + expect(app.store.getState().document.assets.some((asset) => asset.id === "asset-new")).toBe(true); + expect(app.store.getState().document.artboards[0]?.layers.some((layer) => layer.id === "layer-new")).toBe(true); + expect(app.store.getState().editor.generation.candidates).toHaveLength(0); + }); +}); + +function dependencies(overrides: Partial): GenerationWorkflowDependencies { + return { + runGenerate: async () => undefined, + runGenerateFromCandidate: async () => undefined, + createMaskedPixelReplacementSource: async () => "replacement", + createRefinementMask: async () => "mask", + loadGenerationResources: async () => undefined, + createId: () => crypto.randomUUID(), + ...overrides, + }; +} + +function candidate(artboardId: string): GenerationCandidate { + return { + id: "candidate", + source: "generated", + mimeType: "image/png", + intrinsicSize: { w: 64, h: 64 }, + mode: "text-to-image", + settings: initialToolState.generate, + seed: 1, + width: 64, + height: 64, + placement: { + artboardId, + layerName: "Generated", + transform: { position: { x: 0, y: 0 }, scale: { x: 1, y: 1 }, rotation: 0 }, + }, + }; +} + +function createTestApp() { + const state = createInitialAppState("Test"); + state.document.artboards.push({ + id: "artboard", + name: "Artboard", + bounds: { x: 0, y: 0, w: 100, h: 100 }, + backgroundColor: "transparent", + visible: true, + locked: false, + layers: [], + }); + const registry = createCommandRegistry([...documentCommands, ...toolCommands, ...generationCommands]); + return { store: createAppStore(state, registry) }; +} diff --git a/operations/generation/workflow.ts b/operations/generation/workflow.ts new file mode 100644 index 0000000..f2ff48a --- /dev/null +++ b/operations/generation/workflow.ts @@ -0,0 +1,125 @@ +import { commandIds } from "@commands/ids"; +import type { GenerationCandidate, GenerationJobKind } from "@editor/state"; +import type { AppStore } from "@editor/store"; +import type { GenerateSettings } from "@editor/tools"; +import { createRefinementMask } from "@operations/masks/rasterActions"; +import { createMaskedPixelReplacementSource } from "./candidateActions"; +import { runGenerationJob } from "./generationJob"; +import { loadGenerationResources } from "./loadResources"; +import { runGenerate, runGenerateFromCandidate } from "./runGenerate"; +import { checkGenerationPreconditions } from "./preconditions"; + +export type GenerationWorkflow = ReturnType; + +export type GenerationWorkflowDependencies = { + runGenerate: typeof runGenerate; + runGenerateFromCandidate: typeof runGenerateFromCandidate; + createMaskedPixelReplacementSource: typeof createMaskedPixelReplacementSource; + createRefinementMask: typeof createRefinementMask; + loadGenerationResources: typeof loadGenerationResources; + createId(): string; +}; + +const defaultDependencies: GenerationWorkflowDependencies = { + runGenerate, + runGenerateFromCandidate, + createMaskedPixelReplacementSource, + createRefinementMask, + loadGenerationResources, + createId: () => crypto.randomUUID(), +}; + +export function createGenerationWorkflow(store: AppStore, dependencies: GenerationWorkflowDependencies = defaultDependencies) { + const job = (kind: GenerationJobKind, label: string, task: () => Promise) => + runGenerationJob({ kind, label, dispatch: store.dispatch, task }); + + return { + precondition: () => { + const state = store.getState(); + return checkGenerationPreconditions(state.document, state.editor.selection, state.editor.tools.generate); + }, + + loadResources: () => dependencies.loadGenerationResources(store), + + generate: () => job("generate", "Generating", async () => { + const state = store.getState(); + await dependencies.runGenerate({ + document: state.document, + selection: state.editor.selection, + viewport: state.editor.viewport, + settings: state.editor.tools.generate, + dispatch: store.dispatch, + }); + }), + + regenerate: (candidateId: string, settings?: GenerateSettings, label = "Regenerate") => + job("regenerate", label, async () => { + const candidate = findCandidate(store, candidateId); + const nextSettings = settings ?? candidate.settings; + store.dispatch(commandIds.toolSetGenerateSettings, nextSettings); + await dependencies.runGenerateFromCandidate({ candidate, settings: nextSettings, dispatch: store.dispatch }); + }), + + applyCandidateAsLayer: (candidateId: string) => { + store.dispatch(commandIds.generationApplyCandidateAsLayer, { + candidateId, + assetId: dependencies.createId(), + layerId: dependencies.createId(), + }); + }, + + applyCandidateAsRefinementLayer: (candidateId: string) => + job("refine", "Adding refinement mask", async () => { + const candidate = findCandidate(store, candidateId); + const layerId = dependencies.createId(); + const maskAssetId = dependencies.createId(); + const width = Math.max(1, Math.round(candidate.intrinsicSize.w)); + const height = Math.max(1, Math.round(candidate.intrinsicSize.h)); + const source = await dependencies.createRefinementMask(width, height); + + store.dispatch(commandIds.generationApplyCandidateAsLayer, { + candidateId, + assetId: dependencies.createId(), + layerId, + }); + store.dispatch(commandIds.documentAddLayerMask, { + layerId, + asset: { + id: maskAssetId, + name: `${candidate.placement.layerName} refinement mask`, + mimeType: "image/png", + source, + intrinsicSize: { w: width, h: height }, + }, + maskLayer: { + id: dependencies.createId(), + type: "raster", + name: `${candidate.placement.layerName} refinement mask`, + visible: true, + locked: false, + opacity: 1, + assetId: maskAssetId, + transform: { + position: { ...candidate.placement.transform.position }, + scale: { ...candidate.placement.transform.scale }, + rotation: candidate.placement.transform.rotation, + }, + }, + }); + store.dispatch(commandIds.toolSetActive, { tool: "eraser" }); + }), + + replaceCandidatePixels: (candidateId: string) => + job("replace", "Replacing pixels", async () => { + const candidate = findCandidate(store, candidateId); + const source = await dependencies.createMaskedPixelReplacementSource(store.getState().document, candidate); + store.dispatch(commandIds.generationReplaceCandidatePixels, { candidateId, source, mimeType: "image/png" }); + }), + }; +} + +function findCandidate(store: AppStore, candidateId: string): GenerationCandidate { + const candidate = store.getState().editor.generation.candidates.find((item) => item.id === candidateId); + if (!candidate) throw new Error("This generation candidate is no longer available."); + return candidate; +} diff --git a/view/App.tsx b/view/App.tsx index 20ce645..891acd3 100644 --- a/view/App.tsx +++ b/view/App.tsx @@ -18,7 +18,6 @@ import { shallowEqual, useAppState } from "./useAppState"; import { downloadArtboardPng } from "@operations/export/downloadArtboard"; import { useImageImport } from "./useImageImport"; import { useViewportActivityIsland } from "./useViewportActivityIsland"; -import { loadGenerationResources } from "@operations/generation/loadResources"; import { useProjectLifecycle } from "./useProjectLifecycle"; import "./index.css"; @@ -38,8 +37,8 @@ export function App({ app }: AppProps) { const layersOpen = workspace.panel === "layers"; useEffect(() => { - if (generateOpen) void loadGenerationResources(app.store); - }, [app.store, generateOpen]); + if (generateOpen) void app.workflows.generation.loadResources(); + }, [app.workflows.generation, generateOpen]); const openGenerate = useCallback(() => { app.store.dispatch(commandIds.workspaceSetPanel, { panel: "generate" }); @@ -215,6 +214,7 @@ export function App({ app }: AppProps) { transformBounds={viewportActivityIsland.visible ? undefined : transformBounds} transformTarget={viewportActivityIsland.visible ? undefined : transformTarget} dispatch={app.store.dispatch} + generationWorkflow={app.workflows.generation} />
diff --git a/view/BottomControlsIsland.tsx b/view/BottomControlsIsland.tsx index d1886a3..ab132f9 100644 --- a/view/BottomControlsIsland.tsx +++ b/view/BottomControlsIsland.tsx @@ -11,6 +11,7 @@ import { TransformControls } from "./bottom-controls/TransformControls"; import { ZoomControls } from "./bottom-controls/ZoomControls"; import type { Rect } from "@core/geometry"; import type { TransformTarget } from "@editor/transform"; +import type { GenerationWorkflow } from "@operations/generation/workflow"; export type BottomControlsAction = "pan" | "zoom"; export type BottomControlsIslandProps = { @@ -31,9 +32,10 @@ export type BottomControlsIslandProps = { transformTarget?: TransformTarget; brushHint?: string; dispatch: AppStore["dispatch"]; + generationWorkflow: GenerationWorkflow; }; -export function BottomControlsIsland({ document, selection, viewport, visible, action, activeTool, brushSettings, generateSettings, generation, chromaKeySettings, magicWandSettings, editingMask = false, maskViewMode = "composite", transformBounds, transformTarget, brushHint, dispatch }: BottomControlsIslandProps) { +export function BottomControlsIsland({ document, selection, viewport, visible, action, activeTool, brushSettings, generateSettings, generation, chromaKeySettings, magicWandSettings, editingMask = false, maskViewMode = "composite", transformBounds, transformTarget, brushHint, dispatch, generationWorkflow }: BottomControlsIslandProps) { const zoomPercent = Math.round(viewport.zoom * 100); const x = Math.round(viewport.center.x); const y = Math.round(viewport.center.y); @@ -46,7 +48,7 @@ export function BottomControlsIsland({ document, selection, viewport, visible, a }`} > {activeTool === "generate" ? ( - + ) : (activeTool === "brush" || activeTool === "eraser") && brushHint ? ( ) : activeTool === "brush" || activeTool === "eraser" ? ( diff --git a/view/bottom-controls/GenerateActionControls.tsx b/view/bottom-controls/GenerateActionControls.tsx index 448c95f..9acaf81 100644 --- a/view/bottom-controls/GenerateActionControls.tsx +++ b/view/bottom-controls/GenerateActionControls.tsx @@ -1,29 +1,22 @@ import { commandIds } from "@commands/ids"; -import type { ImageDocument } from "@core/document"; -import type { GenerationCandidate, GenerationCompareMode, GenerationState, SelectionState, ViewportState } from "@editor/state"; +import type { GenerationCandidate, GenerationCompareMode, GenerationState } from "@editor/state"; import type { GenerateSettings } from "@editor/tools"; import type { AppStore } from "@editor/store"; -import { createMaskedPixelReplacementSource } from "@operations/generation/candidateActions"; -import { runGenerate, runGenerateFromCandidate } from "@operations/generation/runGenerate"; -import { runGenerationJob } from "@operations/generation/generationJob"; +import type { GenerationWorkflow } from "@operations/generation/workflow"; import { currentGenerationJob, GenerationJobStatus } from "../GenerationJobStatus"; -import { createRefinementMask } from "@operations/masks/rasterActions"; -import { checkGenerationPreconditions } from "@operations/generation/preconditions"; export type GenerateActionControlsProps = { - document: ImageDocument; - selection: SelectionState; - viewport: ViewportState; settings: GenerateSettings; generation: GenerationState; dispatch: AppStore["dispatch"]; + workflow: GenerationWorkflow; }; -export function GenerateActionControls({ document, selection, viewport, settings, generation, dispatch }: GenerateActionControlsProps) { +export function GenerateActionControls({ settings, generation, dispatch, workflow }: GenerateActionControlsProps) { const job = currentGenerationJob(generation); const busy = job?.status === "running"; const candidate = selectedCandidate(generation); - const precondition = checkGenerationPreconditions(document, selection, settings); + const precondition = workflow.precondition(); const canGenerate = precondition.ready && !busy; const preconditionMessage = precondition.ready ? undefined : precondition.message; @@ -35,7 +28,7 @@ export function GenerateActionControls({ document, selection, viewport, settings className="h-12 rounded-full bg-white px-7 text-base font-semibold !text-black transition hover:bg-white/90 focus:outline-none focus-visible:ring-2 focus-visible:ring-white/40 disabled:pointer-events-none disabled:opacity-35" title={job?.status === "failed" ? job.error : preconditionMessage ?? "Generate with ComfyUI"} onClick={() => { - void runGenerationJob({ kind: "generate", label: "Generating", dispatch, task: () => runGenerate({ document, selection, viewport, settings, dispatch }) }); + void workflow.generate(); }} > {busy && job?.kind === "generate" ? "Generating..." : "Generate"} @@ -46,12 +39,12 @@ export function GenerateActionControls({ document, selection, viewport, settings <> ) : null} @@ -82,23 +75,22 @@ function CandidatePicker({ generation, dispatch }: { generation: GenerationState } function CandidateControls({ - document, candidate, compareMode, settings, busy, dispatch, + workflow, }: { - document: ImageDocument; candidate: GenerationCandidate; compareMode: GenerationCompareMode; settings: GenerateSettings; busy: boolean; dispatch: AppStore["dispatch"]; + workflow: GenerationWorkflow; }) { const rerun = (label: string, nextSettings: GenerateSettings) => { - dispatch(commandIds.toolSetGenerateSettings, nextSettings); - void runGenerationJob({ kind: "regenerate", label, dispatch, task: () => runGenerateFromCandidate({ candidate, settings: nextSettings, dispatch }) }); + void workflow.regenerate(candidate.id, nextSettings, label); }; const disabled = busy; @@ -116,13 +108,13 @@ function CandidateControls({ /> rerun("Reuse seed", { ...candidate.settings, seed: candidate.seed })} /> rerun("New seed", { ...candidate.settings, seed: -1 })} /> - applyCandidateAsLayer(candidate, dispatch)} /> + workflow.applyCandidateAsLayer(candidate.id)} /> { - void runGenerationJob({ kind: "refine", label: "Adding refinement mask", dispatch, task: () => applyCandidateAsRefinementLayer(candidate, dispatch) }); + void workflow.applyCandidateAsRefinementLayer(candidate.id); }} /> { - void runGenerationJob({ kind: "replace", label: "Replacing pixels", dispatch, task: async () => { - const source = await createMaskedPixelReplacementSource(document, candidate); - dispatch(commandIds.generationReplaceCandidatePixels, { candidateId: candidate.id, source, mimeType: "image/png" }); - } }); + void workflow.replaceCandidatePixels(candidate.id); }} /> candidate.id === generation.selectedCandidateId) ?? generation.candidates[0]; } - -function applyCandidateAsLayer(candidate: GenerationCandidate, dispatch: AppStore["dispatch"]) { - applyCandidateAsLayerWithIds(candidate, { layerId: crypto.randomUUID(), assetId: crypto.randomUUID() }, dispatch); -} - -async function applyCandidateAsRefinementLayer(candidate: GenerationCandidate, dispatch: AppStore["dispatch"]) { - const layerId = crypto.randomUUID(); - const maskLayerId = crypto.randomUUID(); - const maskAssetId = crypto.randomUUID(); - applyCandidateAsLayerWithIds(candidate, { layerId, assetId: crypto.randomUUID() }, dispatch); - const width = Math.max(1, Math.round(candidate.intrinsicSize.w)); - const height = Math.max(1, Math.round(candidate.intrinsicSize.h)); - const source = await createRefinementMask(width, height); - - dispatch(commandIds.documentAddLayerMask, { - layerId, - asset: { - id: maskAssetId, - name: `${candidate.placement.layerName} refinement mask`, - mimeType: "image/png", - source, - intrinsicSize: { w: width, h: height }, - }, - maskLayer: { - id: maskLayerId, - type: "raster", - name: `${candidate.placement.layerName} refinement mask`, - visible: true, - locked: false, - opacity: 1, - assetId: maskAssetId, - transform: { - position: { ...candidate.placement.transform.position }, - scale: { ...candidate.placement.transform.scale }, - rotation: candidate.placement.transform.rotation, - }, - }, - }); - dispatch(commandIds.toolSetActive, { tool: "eraser" }); -} - -function applyCandidateAsLayerWithIds(candidate: GenerationCandidate, ids: { layerId: string; assetId: string }, dispatch: AppStore["dispatch"]) { - dispatch(commandIds.generationApplyCandidateAsLayer, { - candidateId: candidate.id, - assetId: ids.assetId, - layerId: ids.layerId, - }); -} diff --git a/view/bottom-controls/GenerateControls.tsx b/view/bottom-controls/GenerateControls.tsx index a4f3cd2..a3bf8ea 100644 --- a/view/bottom-controls/GenerateControls.tsx +++ b/view/bottom-controls/GenerateControls.tsx @@ -2,9 +2,9 @@ import { useEffect, useRef, useState, type RefObject } from "react"; import { CaretDown, CaretUp } from "@phosphor-icons/react"; import { commandIds } from "@commands/ids"; import type { AppStore } from "@editor/store"; -import type { GenerationOptions, GenerationResourcesState } from "@editor/state"; -import { generateArchitectureDefaults } from "@editor/tools"; -import type { GenerateArchitecture, GenerateMode, GenerateModel, GenerateSettings } from "@editor/tools"; +import type { GenerationResourcesState } from "@editor/state"; +import type { GenerateArchitecture, GenerateMode, GenerateSettings } from "@editor/tools"; +import { resolveGenerationModeOptions, resolveGenerationModelOptions, resolveGenerationStringOptions, resolveGenerationSupportOptions } from "@operations/generation/options"; import { BottomControlSelectMenu, type BottomControlSelectOption } from "./SelectMenu"; import { BottomControlSlider } from "./Slider"; @@ -55,11 +55,11 @@ export function GenerateControls({ settings, resources, dispatch }: GenerateCont const [inpaintOpen, setInpaintOpen] = useState(false); const [sizeOpen, setSizeOpen] = useState(false); const sizeRef = useRef(null); - const modelOptions = resolveModelOptions(settings, comfyOptions); - const supportOptions = resolveSupportOptions(settings, comfyOptions); - const samplerOptions = resolveStringOptions(comfyOptions?.samplers, settings.sampler); - const schedulerOptions = resolveStringOptions(comfyOptions?.schedulers, settings.scheduler); - const modeOptions = resolveModeOptions(settings, comfyOptions); + const modelOptions = resolveGenerationModelOptions(settings, comfyOptions); + const supportOptions = resolveGenerationSupportOptions(settings, comfyOptions); + const samplerOptions = resolveGenerationStringOptions(comfyOptions?.samplers, settings.sampler); + const schedulerOptions = resolveGenerationStringOptions(comfyOptions?.schedulers, settings.scheduler); + const modeOptions = resolveGenerationModeOptions(settings, comfyOptions, modes); useEffect(() => { if (!sizeOpen) return; @@ -197,38 +197,6 @@ export function GenerateControls({ settings, resources, dispatch }: GenerateCont ); } -function resolveModelOptions(settings: GenerateSettings, comfyOptions: GenerationOptions | undefined): readonly BottomControlSelectOption[] { - const architecture = comfyOptions?.architectures?.find((option) => option.value === settings.architecture); - const models = architecture?.models ?? (settings.architecture === "sdxl" ? comfyOptions?.models : undefined) ?? []; - const fallbackModel = architecture?.defaultModel ?? generateArchitectureDefaults[settings.architecture].model; - const values = unique(["auto", ...models, ...(models.length === 0 && fallbackModel !== "auto" ? [fallbackModel] : []), settings.model]); - return values.map((model) => ({ value: model, label: model === "auto" ? "Auto" : model })); -} - -function resolveSupportOptions(settings: GenerateSettings, comfyOptions: GenerationOptions | undefined) { - const defaults = generateArchitectureDefaults[settings.architecture]; - return { - textEncoders: resolveStringOptions([...(comfyOptions?.textEncoders ?? []), defaults.textEncoder].filter((value) => value !== "auto"), settings.textEncoder), - vaes: resolveStringOptions([...(comfyOptions?.vaes ?? []), defaults.vae].filter((value) => value !== "auto"), settings.vae), - }; -} - -function resolveStringOptions(values: string[] | undefined, current: string): readonly BottomControlSelectOption[] { - return unique([...(values ?? []), current]).map((value) => ({ value, label: value })); -} - -function resolveModeOptions(settings: GenerateSettings, comfyOptions: GenerationOptions | undefined): readonly BottomControlSelectOption[] { - const architecture = comfyOptions?.architectures?.find((option) => option.value === settings.architecture); - const supportedModes = architecture?.supportedModes?.length ? architecture.supportedModes : generateArchitectureDefaults[settings.architecture].supportedModes; - const availableModes = modes.filter((mode) => supportedModes.includes(mode.value)); - if (availableModes.some((mode) => mode.value === settings.mode)) return availableModes; - return [modes.find((mode) => mode.value === settings.mode), ...availableModes].filter((mode): mode is BottomControlSelectOption => Boolean(mode)); -} - -function unique(values: T[]): T[] { - return Array.from(new Set(values)); -} - function SectionTitle({ title }: { title: string }) { return
{title}
; }