feat: implement generation workflow with resource loading and candidate management

This commit is contained in:
syntaxbullet
2026-07-10 23:54:52 +02:00
parent 9697075d29
commit 51a54dbdb2
9 changed files with 308 additions and 117 deletions

View File

@@ -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"]);
});
});

View File

@@ -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<TValue extends string> = { value: TValue; label: string };
export function resolveGenerationModelOptions(settings: GenerateSettings, options: GenerationOptions | undefined): readonly GenerationSelectOption<GenerateModel>[] {
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<string>[] {
return unique([...(values ?? []), current]).map((value) => ({ value, label: value }));
}
export function resolveGenerationModeOptions(
settings: GenerateSettings,
options: GenerationOptions | undefined,
modes: readonly GenerationSelectOption<GenerateMode>[],
): readonly GenerationSelectOption<GenerateMode>[] {
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<GenerateMode> => Boolean(mode));
}
function unique<T>(values: T[]): T[] {
return Array.from(new Set(values));
}

View File

@@ -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>): 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) };
}

View File

@@ -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<typeof createGenerationWorkflow>;
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<void>) =>
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;
}