178 lines
7.4 KiB
TypeScript
178 lines
7.4 KiB
TypeScript
import { commandIds } from "@commands/ids";
|
|
import type { Layer } from "@core/layer";
|
|
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) {
|
|
let activeController: AbortController | undefined;
|
|
const job = async (kind: GenerationJobKind, label: string, task: (signal: AbortSignal) => Promise<void>) => {
|
|
if (store.getState().editor.generation.jobs.some((candidate) => candidate.status === "running")) return;
|
|
const controller = new AbortController();
|
|
activeController = controller;
|
|
try {
|
|
await runGenerationJob({ kind, label, dispatch: store.dispatch, signal: controller.signal, task });
|
|
} finally {
|
|
if (activeController === controller) activeController = undefined;
|
|
}
|
|
};
|
|
|
|
return {
|
|
precondition: () => {
|
|
const state = store.getState();
|
|
return checkGenerationPreconditions(state.document, state.editor.selection, state.editor.tools.generate);
|
|
},
|
|
|
|
loadResources: () => dependencies.loadGenerationResources(store),
|
|
|
|
prepareInpaintMask: async () => {
|
|
const state = store.getState();
|
|
const layerId = state.editor.selection.layerIds.length === 1 ? state.editor.selection.layerIds[0] : undefined;
|
|
if (!layerId) return;
|
|
const artboard = state.document.artboards.find((candidate) => candidate.id === state.editor.selection.artboardId);
|
|
const layer = artboard ? findLayer(artboard.layers, layerId) : undefined;
|
|
if (!layer || layer.type === "group" || layer.type === "adjustment") return;
|
|
const asset = state.document.assets.find((candidate) => candidate.id === layer.assetId);
|
|
if (!asset) return;
|
|
const source = await dependencies.createRefinementMask(asset.intrinsicSize.w, asset.intrinsicSize.h);
|
|
const maskAssetId = dependencies.createId();
|
|
const maskLayerId = dependencies.createId();
|
|
store.dispatch(commandIds.documentAddLayerMask, {
|
|
layerId,
|
|
asset: { id: maskAssetId, name: `${layer.name} AI edit mask`, mimeType: "image/png", source, intrinsicSize: { ...asset.intrinsicSize } },
|
|
maskLayer: {
|
|
id: maskLayerId,
|
|
type: "raster",
|
|
name: `${layer.name} AI edit mask`,
|
|
visible: true,
|
|
locked: false,
|
|
opacity: 1,
|
|
assetId: maskAssetId,
|
|
transform: { position: { ...layer.transform.position }, scale: { ...layer.transform.scale }, rotation: layer.transform.rotation },
|
|
},
|
|
});
|
|
},
|
|
|
|
generate: () => job("generate", "Generating", async (signal) => {
|
|
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,
|
|
signal,
|
|
});
|
|
}),
|
|
|
|
regenerate: (candidateId: string, settings?: GenerateSettings, label = "Regenerate") =>
|
|
job("regenerate", label, async (signal) => {
|
|
const candidate = findCandidate(store, candidateId);
|
|
const nextSettings = settings ?? candidate.settings;
|
|
store.dispatch(commandIds.toolSetGenerateSettings, nextSettings);
|
|
await dependencies.runGenerateFromCandidate({ candidate, settings: nextSettings, dispatch: store.dispatch, signal });
|
|
}),
|
|
|
|
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" });
|
|
}),
|
|
|
|
cancel: () => activeController?.abort(),
|
|
};
|
|
}
|
|
|
|
function findLayer(layers: readonly Layer[], layerId: string): Layer | undefined {
|
|
for (const layer of layers) {
|
|
if (layer.id === layerId) return layer;
|
|
if (layer.type === "group") {
|
|
const child = findLayer(layer.children, layerId);
|
|
if (child) return child;
|
|
}
|
|
}
|
|
return undefined;
|
|
}
|
|
|
|
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;
|
|
}
|