diff --git a/pages/index.vue b/pages/index.vue index be3ace3..747f87f 100644 --- a/pages/index.vue +++ b/pages/index.vue @@ -522,11 +522,12 @@ Remove - @@ -1417,18 +1418,15 @@ {{ lastFrameLoading ? 'Reading last frame…' : 'Last frame preview unavailable' }}

-
- Extension prompt - -
+ {{ watchingLiveVideo() || videoBusy ? 'Queue extension' : 'Generate Extension' }} @@ -2364,7 +2362,7 @@ interface RetryDraft { fps?: number samplerName?: string scheduler?: string - extensions?: { prompt: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] + extensions?: { prompt: string; promptPre?: string; promptPost?: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] workflow?: string globalLocks?: string permanenceRefs?: PermanenceRef[] @@ -2378,6 +2376,8 @@ interface RetryDraft { interface QueuedExtension { id: string prompt: string + promptPre: string + promptPost: string duration: number loraName?: string loraStack?: LoraStackItem[] @@ -2602,6 +2602,8 @@ const currentClipId = ref('') const extendOpen = ref(false) const extendClipId = ref('') const extendPrompt = ref('') +const extendPromptPre = ref('') +const extendPromptPost = ref('') const extendDuration = ref(5) const lastFrameUrl = ref('') const lastFrameLoading = ref(false) @@ -2998,7 +3000,7 @@ const canExtend = computed(() => Boolean( const extendSaveName = computed(() => nextClipPartName(currentClip.value?.name || currentClip.value?.prompt || '')) const queueReady = computed(() => { if (shotScriptMode.value) return parsedShots.value.length > 0 && parsedShots.value.every(shot => shot.prompt.trim()) - return extensionQueue.value.every(item => item.prompt.trim()) + return extensionQueue.value.every(item => composePromptParts(item.promptPre, item.prompt, item.promptPost)) }) const canUseGeneratedStill = computed(() => Boolean( editResultUrl.value && !editLockedSave.value && !editAwaitingReveal.value @@ -3204,10 +3206,10 @@ const plannedChain = computed(() => { })) } return [ - { label: 'Initial', prompt: prompt.value, duration: duration.value }, + { label: 'Initial', prompt: composePromptParts(promptPre.value, prompt.value, promptPost.value), duration: duration.value }, ...extensionQueue.value.map((item, index) => ({ label: `Extension ${index + 1}`, - prompt: item.prompt, + prompt: composePromptParts(item.promptPre, item.prompt, item.promptPost), duration: item.duration })) ] @@ -4128,6 +4130,26 @@ function applyPromptRestore(saved: string, stored?: { prompt?: string; promptPre promptPost.value = parts.post } +function queuedExtensionFrom( + item: { prompt?: string; promptPre?: string; promptPost?: string; duration?: number; loraName?: string; loraStack?: LoraStackItem[] }, + duration: number, + loraStack: LoraStackItem[] +): QueuedExtension { + const parts = restorePromptParts(item.prompt || '', { + pre: item.promptPre, + prompt: item.prompt, + post: item.promptPost + }) + return { + id: crypto.randomUUID(), + prompt: parts.prompt, + promptPre: parts.pre, + promptPost: parts.post, + duration, + loraStack + } +} + function insertKeepExtend(item: KeepPrompt) { extendPrompt.value = insertKeepSnippet(extendPrompt.value, item.text) } @@ -5155,7 +5177,7 @@ async function restoreStudioJob(id: string) { hideThumbnail?: boolean hideInput?: boolean referenceStillIds?: Array - extensions?: { prompt: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] + extensions?: { prompt: string; promptPre?: string; promptPost?: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] queueAutoRun?: boolean promptPre?: string promptPost?: string @@ -5174,10 +5196,14 @@ async function restoreStudioJob(id: string) { outputFocus.value = 'video' closeExtend() const extensions = payload.extensions || [] + const queuedAsCards = extensions.some(item => + Object.prototype.hasOwnProperty.call(item, 'promptPre') + || Object.prototype.hasOwnProperty.call(item, 'promptPost') + ) applyVideoWorkflow(payload.workflow) - applyPromptRestore(extensions.length - ? [payload.prompt.trim(), ...extensions.map((item, index) => `shot ${index + 2}\n${item.prompt.trim()}`)].join('\n\n') - : (payload.prompt || ''), payload) + applyPromptRestore(queuedAsCards || !extensions.length + ? (payload.prompt || '') + : [payload.prompt.trim(), ...extensions.map((item, index) => `shot ${index + 2}\n${item.prompt.trim()}`)].join('\n\n'), payload) clipName.value = payload.name || '' seedInput.value = String(payload.seed || '') restoreAdvanced({ @@ -5198,14 +5224,13 @@ async function restoreStudioJob(id: string) { browseFolderId.value = payload.folderId } if (typeof payload.sound === 'boolean') withSound.value = payload.sound - shotScriptMode.value = extensions.length > 0 + shotScriptMode.value = extensions.length > 0 && !queuedAsCards identityRefs.value.forEach((_, index) => clearIdentityRef(index)) - extensionQueue.value = extensions.map(item => ({ - id: crypto.randomUUID(), - prompt: item.prompt, - duration: clampDuration(Number(item.duration)), - loraStack: readLoraStack(item) - })) + extensionQueue.value = extensions.map(item => queuedExtensionFrom( + item, + clampDuration(Number(item.duration)), + readLoraStack(item) + )) queueAutoRun.value = payload.queueAutoRun === true globalLocks.value = payload.globalLocks || '' familyPermanenceRefs.value = payload.permanenceRefs || [] @@ -5590,12 +5615,11 @@ async function rerun(item: LibraryClip, collection = false) { restoreAll ? ordered : undefined ) extensionQueue.value = restoreAll - ? parts.slice(1).map(part => ({ - id: crypto.randomUUID(), - prompt: part.prompt, - duration: clipGenerateDuration(part), - loraStack: loraStacksEqual(readLoraStack(part), readLoraStack(ordered[0] || target)) ? [] : readLoraStack(part) - })) + ? parts.slice(1).map(part => queuedExtensionFrom( + part, + clipGenerateDuration(part), + loraStacksEqual(readLoraStack(part), readLoraStack(ordered[0] || target)) ? [] : readLoraStack(part) + )) : [] try { const initialTextToVideo = textToVideo.value && !target.parentClipId && !(target.chainIndex) @@ -5636,12 +5660,11 @@ async function loadDraft(draft: RetryDraft) { }) shotPermanenceRefs.value = draftShots applyLoraSelection(draft, draft.shotLoraStacks || draft.shotLoras) - extensionQueue.value = (draft.extensions || []).map(item => ({ - id: crypto.randomUUID(), - prompt: String(item.prompt || ''), - duration: clampDuration(Number(item.duration)), - loraStack: readLoraStack(item) - })) + extensionQueue.value = (draft.extensions || []).map(item => queuedExtensionFrom( + item, + clampDuration(Number(item.duration)), + readLoraStack(item) + )) try { const blob = await $fetch(`/api/library/drafts/${draft.id}/still`, { responseType: 'blob' }) const filename = draft.stillFilename || 'held-still.png' @@ -5715,6 +5738,9 @@ function closeExtend() { lastFrameGen += 1 extendOpen.value = false extendClipId.value = '' + extendPromptPre.value = '' + extendPrompt.value = '' + extendPromptPost.value = '' lastFrameLoading.value = false lastFrameRevealed.value = false if (lastFrameObjectUrl) { @@ -5760,8 +5786,14 @@ async function openExtend(item?: LibraryClip) { currentClipId.value = id extendClipId.value = id const clip = clips.value.find(entry => entry.id === id) - if (clip) applyPromptRestore(clip.prompt, clip) - extendPrompt.value = clip?.prompt || prompt.value + const parts = restorePromptParts(clip?.prompt || prompt.value, { + pre: clip ? (clip.promptPre || '') : promptPre.value, + prompt: clip?.prompt || prompt.value, + post: clip ? (clip.promptPost || '') : promptPost.value + }) + extendPromptPre.value = parts.pre + extendPrompt.value = parts.prompt + extendPromptPost.value = parts.post extendDuration.value = clampDuration(duration.value) extendLoraStack.value = readLoraStack(clip) lastFrameRevealed.value = false @@ -5776,7 +5808,7 @@ async function openExtend(item?: LibraryClip) { async function generateExtension() { const id = extendClipId.value || clipIdFromPlayer() - if (!extendPrompt.value.trim()) return + if (!composePromptParts(extendPromptPre.value, extendPrompt.value, extendPromptPost.value)) return if (!id) { toast('Load a generated clip before extending it.') return @@ -5799,8 +5831,8 @@ async function generateExtension() { body: { clipId: id, prompt: extendPrompt.value.trim(), - promptPre: promptPre.value.trim() || undefined, - promptPost: promptPost.value.trim() || undefined, + promptPre: extendPromptPre.value, + promptPost: extendPromptPost.value, duration: extendDuration.value, loraStack: extendLoraStack.value } @@ -6475,6 +6507,8 @@ function queueExtension() { extensionQueue.value.push({ id: crypto.randomUUID(), prompt: prompt.value.trim(), + promptPre: promptPre.value, + promptPost: promptPost.value, duration: duration.value, loraStack: [] }) @@ -7012,7 +7046,13 @@ async function generate() { const initialPrompt = (shots?.[0]?.prompt || prompt.value).trim() const queued = shots ? shots.slice(1).map(shot => ({ prompt: shot.prompt.trim(), duration: duration.value, ...persistLoraFields(shotLoraStacks.value[shot.n] || []) })) - : extensionQueue.value.map(item => ({ prompt: item.prompt.trim(), duration: item.duration, ...persistLoraFields(item.loraStack || []) })) + : extensionQueue.value.map(item => ({ + prompt: item.prompt.trim(), + promptPre: item.promptPre, + promptPost: item.promptPost, + duration: item.duration, + ...persistLoraFields(item.loraStack || []) + })) const hideOut = hideThumbnail.value const folderLockedPref = false const keepLiveOutput = watchingLiveVideo() diff --git a/server/api/extend.post.ts b/server/api/extend.post.ts index 01f27af..5308c6b 100644 --- a/server/api/extend.post.ts +++ b/server/api/extend.post.ts @@ -4,6 +4,7 @@ import { persistLoraFields, parsePostedLoraStack, listStudioLoras } from '~/serv import { readLoraStack } from '~/utils/loras' import { parseExtendDuration } from '~/server/utils/extendChain' import { defaultVideoSteps, isLtxWorkflow, LTX_DISABLED_MESSAGE, ltxWorkflowEnabled, parseVideoWorkflow } from '~/utils/videoModels' +import { composePromptParts } from '~/utils/promptParts' export default defineEventHandler(async (event) => { const body = await readBody<{ clipId?: string; prompt?: string; promptPre?: string; promptPost?: string; duration?: number; lora?: unknown; loraStack?: unknown }>(event).catch(() => ({})) @@ -12,9 +13,6 @@ export default defineEventHandler(async (event) => { if (!clipId) { throw createError({ statusCode: 400, statusMessage: 'A clip is required to extend' }) } - if (!prompt) { - throw createError({ statusCode: 400, statusMessage: 'An extension prompt is required' }) - } const { owner } = assertLibraryOwner(event) const source = getClip(owner, clipId) @@ -27,8 +25,15 @@ export default defineEventHandler(async (event) => { const destFolder = publicLibrary(event).folders.find(folder => folder.id === source.folderId) const folderLocked = Boolean(destFolder?.protected && !destFolder.unlocked) const durationSeconds = parseExtendDuration(body?.duration) - const promptPre = String(body?.promptPre ?? source.promptPre ?? '').trim() - const promptPost = String(body?.promptPost ?? source.promptPost ?? '').trim() + const promptPre = body && Object.prototype.hasOwnProperty.call(body, 'promptPre') + ? String(body.promptPre || '').trim() + : String(source.promptPre || '').trim() + const promptPost = body && Object.prototype.hasOwnProperty.call(body, 'promptPost') + ? String(body.promptPost || '').trim() + : String(source.promptPost || '').trim() + if (!composePromptParts(promptPre, prompt, promptPost)) { + throw createError({ statusCode: 400, statusMessage: 'An extension prompt is required' }) + } const workflow = parseVideoWorkflow(source.workflow) if (isLtxWorkflow(workflow) && !ltxWorkflowEnabled()) { throw createError({ statusCode: 400, statusMessage: LTX_DISABLED_MESSAGE }) @@ -47,8 +52,8 @@ export default defineEventHandler(async (event) => { payload: { prompt, promptMid: prompt, - promptPre: promptPre || undefined, - promptPost: promptPost || undefined, + promptPre, + promptPost, name: nextClipPartName(clipTitle(source)), folderId: source.folderId, aspect: source.aspect || '16:9', diff --git a/server/api/generate.post.ts b/server/api/generate.post.ts index 6a069f5..df1b05f 100644 --- a/server/api/generate.post.ts +++ b/server/api/generate.post.ts @@ -1,11 +1,11 @@ import { existsSync, readFileSync } from 'node:fs' import { addStudioJob, kickStudioQueue, listStudioJobs, videoJobsBusy } from '~/server/utils/studioQueue' import { listStudioLoras, parsePostedLoraStack, parseShotLoraStacks, persistLoraFields } from '~/server/utils/loras' -import type { LoraStackItem } from '~/utils/loras' import { clampVideoCfg } from '~/utils/generationPresets' import { defaultVideoSteps, isLtxWorkflow, isTextToVideo, LTX_DISABLED_MESSAGE, ltxWorkflowEnabled, parseVideoWorkflow, videoEngineOf } from '~/utils/videoModels' import { allowIdentityRefs, normalizePermanenceRefs, resolveGlobalLocks, type PermanenceRef } from '~/utils/globalLocks' -import { composePromptParts } from '~/utils/promptParts' +import { composePromptParts, persistPromptWrappers } from '~/utils/promptParts' +import type { QueuedExtension } from '~/server/utils/library' function parseDuration(raw: unknown) { const seconds = Number(raw) @@ -18,21 +18,26 @@ function parseExtendDuration(raw: unknown) { } function parseExtensions(raw: string | undefined) { - if (!raw) return [] as { prompt: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] + if (!raw) return [] as QueuedExtension[] try { const parsed = JSON.parse(raw) if (!Array.isArray(parsed)) return [] return parsed - .map((item: { prompt?: unknown; duration?: unknown; permanenceRefs?: unknown; loraName?: unknown; loraStack?: unknown }) => { + .map((item: { prompt?: unknown; promptPre?: unknown; promptPost?: unknown; duration?: unknown; permanenceRefs?: unknown; loraName?: unknown; loraStack?: unknown }) => { const loraStack = parsePostedLoraStack(item?.loraStack ?? item?.loraName, 'video') + const wrappers = persistPromptWrappers({ + ...(item && Object.prototype.hasOwnProperty.call(item, 'promptPre') ? { promptPre: String(item.promptPre ?? '') } : {}), + ...(item && Object.prototype.hasOwnProperty.call(item, 'promptPost') ? { promptPost: String(item.promptPost ?? '') } : {}) + }) return { prompt: String(item?.prompt || '').trim(), duration: parseExtendDuration(item?.duration), permanenceRefs: normalizePermanenceRefs(item?.permanenceRefs), - ...persistLoraFields(loraStack) + ...persistLoraFields(loraStack), + ...wrappers } }) - .filter(item => item.prompt) + .filter(item => composePromptParts(item.promptPre || '', item.prompt, item.promptPost || '')) } catch (error) { if (error && typeof error === 'object' && 'statusCode' in error) throw error return [] diff --git a/server/utils/extendChain.ts b/server/utils/extendChain.ts index 5d2cada..e10429d 100644 --- a/server/utils/extendChain.ts +++ b/server/utils/extendChain.ts @@ -70,13 +70,16 @@ export async function beginExtendFromClip(params: { throw createError({ statusCode: 404, statusMessage: 'Source video is missing' }) } const prompt = String(params.prompt || '').trim() - if (!prompt) { + const durationSeconds = parseExtendDuration(params.duration) + const promptPre = params.promptPre != null + ? String(params.promptPre).trim() + : String(source.promptPre || '').trim() + const promptPost = params.promptPost != null + ? String(params.promptPost).trim() + : String(source.promptPost || '').trim() + if (!composePromptParts(promptPre, prompt, promptPost)) { throw createError({ statusCode: 400, statusMessage: 'An extension prompt is required' }) } - - const durationSeconds = parseExtendDuration(params.duration) - const promptPre = String(params.promptPre ?? source.promptPre ?? '').trim() - const promptPost = String(params.promptPost ?? source.promptPost ?? '').trim() const composedPrompt = composeShotPrompt({ globalLocks: source.globalLocks, prompt: composePromptParts(promptPre, prompt, promptPost), diff --git a/server/utils/jobs.ts b/server/utils/jobs.ts index 381210c..a04d419 100644 --- a/server/utils/jobs.ts +++ b/server/utils/jobs.ts @@ -80,7 +80,7 @@ export interface Job { duration?: number sound?: boolean draftId?: string - extensions?: { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: import('~/utils/loras').LoraStackItem[] }[] + extensions?: import('~/server/utils/library').QueuedExtension[] chainIndex?: number chainStep?: number chainTotal?: number diff --git a/server/utils/library.ts b/server/utils/library.ts index aa873d8..bf86300 100644 --- a/server/utils/library.ts +++ b/server/utils/library.ts @@ -98,6 +98,8 @@ export interface PublicFolder { export interface QueuedExtension { prompt: string + promptPre?: string + promptPost?: string duration: number permanenceRefs?: PermanenceRef[] loraName?: string diff --git a/server/utils/pending.ts b/server/utils/pending.ts index c3dd9ae..48382a9 100644 --- a/server/utils/pending.ts +++ b/server/utils/pending.ts @@ -49,8 +49,8 @@ export interface PendingJob { samplerName?: string scheduler?: string hideInput?: boolean - extensions?: { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: import('~/utils/loras').LoraStackItem[] }[] - remainingExtensions?: { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: import('~/utils/loras').LoraStackItem[] }[] + extensions?: import('~/server/utils/library').QueuedExtension[] + remainingExtensions?: import('~/server/utils/library').QueuedExtension[] currentClipId?: string queueId?: string queueAutoRun?: boolean diff --git a/server/utils/shotQueue.ts b/server/utils/shotQueue.ts index d54aeaf..08a6c39 100644 --- a/server/utils/shotQueue.ts +++ b/server/utils/shotQueue.ts @@ -3,6 +3,8 @@ import { join } from 'node:path' import { getJob, listJobs, type Job } from '~/server/utils/jobs' import { parsePostedLoraStack, persistLoraFields } from '~/server/utils/loras' import type { LoraStackItem } from '~/utils/loras' +import { persistPromptWrappers } from '~/utils/promptParts' +import type { QueuedExtension } from '~/server/utils/library' export type ShotQueueStatus = 'idle' | 'running' | 'paused' | 'complete' | 'error' export type ShotSegmentStatus = 'pending' | 'running' | 'complete' | 'error' @@ -10,6 +12,8 @@ export type ShotSegmentStatus = 'pending' | 'running' | 'complete' | 'error' export interface ShotQueueSegment { index: number prompt: string + promptPre?: string + promptPost?: string duration: number status: ShotSegmentStatus clipId?: string @@ -187,7 +191,7 @@ export async function createShotQueue(params: { loraName?: string loraStack?: LoraStackItem[] initial: { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] } - extensions: { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] + extensions: QueuedExtension[] jobId?: string }): Promise { const now = Date.now() @@ -231,7 +235,8 @@ export async function createShotQueue(params: { duration: item.duration, status: 'pending' as const, permanenceRefs: item.permanenceRefs, - ...persistLoraFields(item.loraStack || item.loraName) + ...persistLoraFields(item.loraStack || item.loraName), + ...persistPromptWrappers(item) })) ] } @@ -473,10 +478,16 @@ export function listOwnersWithQueues() { .filter(owner => existsSync(queuesPath(owner))) } -export function remainingFromQueue(queue: ShotQueue): { prompt: string; duration: number; permanenceRefs?: import('~/utils/globalLocks').PermanenceRef[]; loraName?: string; loraStack?: LoraStackItem[] }[] { +export function remainingFromQueue(queue: ShotQueue): QueuedExtension[] { return queue.segments .filter(segment => segment.index > 0) - .map(segment => ({ prompt: segment.prompt, duration: segment.duration, permanenceRefs: segment.permanenceRefs, ...persistLoraFields(segment.loraStack || segment.loraName) })) + .map(segment => ({ + prompt: segment.prompt, + duration: segment.duration, + permanenceRefs: segment.permanenceRefs, + ...persistLoraFields(segment.loraStack || segment.loraName), + ...persistPromptWrappers(segment) + })) } export function lastCompletedIndex(queue: ShotQueue) { @@ -491,6 +502,7 @@ export function liveSegmentPrompt(queue: ShotQueue, extensionIndex: number) { prompt: segment?.prompt || '', duration: segment?.duration || 5, permanenceRefs: segment?.permanenceRefs, - ...persistLoraFields(segment?.loraStack || segment?.loraName) + ...persistLoraFields(segment?.loraStack || segment?.loraName), + ...persistPromptWrappers(segment) } } diff --git a/server/utils/studioQueue.ts b/server/utils/studioQueue.ts index e7d98de..cdae48c 100644 --- a/server/utils/studioQueue.ts +++ b/server/utils/studioQueue.ts @@ -7,7 +7,7 @@ import { fetchLiveQueue } from '~/server/utils/comfy' import { isLtxWorkflow, isTextToVideo, isXaigenStudio, LTX_DISABLED_MESSAGE, ltxWorkflowEnabled, parseVideoWorkflow, type VideoWorkflowId } from '~/utils/videoModels' import { persistLoraFields } from '~/server/utils/loras' import { resolveLoraStack } from '~/utils/loras' -import { allowIdentityRefs } from '~/utils/globalLocks' +import { allowIdentityRefs, type PermanenceRef } from '~/utils/globalLocks' export type StudioJobStatus = 'waiting' | 'running' | 'held' | 'complete' | 'error' | 'cancelled' export type StudioJobKind = 'video' | 'edit' @@ -39,7 +39,7 @@ export interface StudioJobPayload { hideInput?: boolean folderLocked?: boolean referenceStillIds: Array - extensions: { prompt: string; duration: number; permanenceRefs?: PermanenceRef[]; loraName?: string; loraStack?: import('~/utils/loras').LoraStackItem[] }[] + extensions: import('~/server/utils/library').QueuedExtension[] queueAutoRun: boolean globalLocks?: string permanenceRefs?: PermanenceRef[] diff --git a/server/utils/videoChain.ts b/server/utils/videoChain.ts index 679e568..56607f7 100644 --- a/server/utils/videoChain.ts +++ b/server/utils/videoChain.ts @@ -22,7 +22,8 @@ import { updateShotQueue } from '~/server/utils/shotQueue' import { composeShotPrompt, allowIdentityRefs, type PermanenceRef } from '~/utils/globalLocks' -import { composePromptParts } from '~/utils/promptParts' +import { composePromptParts, resolvePromptWrappers } from '~/utils/promptParts' +import type { QueuedExtension } from '~/server/utils/library' import { persistLoraFields, ensureComfyLoraNames } from '~/server/utils/loras' import { readLoraStack, resolveLoraStack } from '~/utils/loras' import type { LoraStackItem } from '~/utils/loras' @@ -43,7 +44,7 @@ type VideoChainParams = { fps: number samplerName: string scheduler: string - extensions: { prompt: string; duration: number; loraName?: string; loraStack?: LoraStackItem[]; permanenceRefs?: PermanenceRef[] }[] + extensions: QueuedExtension[] workflow: VideoWorkflowId duration: number useIdentityRefs: boolean @@ -323,11 +324,12 @@ export async function continueQueuedExtensions( prompt: live.prompt || extensions[i].prompt, duration: live.duration || extensions[i].duration } + const wrappers = resolvePromptWrappers(live, extensions[i], job.library) const shotStack = resolveLoraStack( readLoraStack(liveQueue).length ? readLoraStack(liveQueue) : (params.loraStack || params.loraName), live.loraStack || live.loraName || extensions[i]?.loraStack || extensions[i]?.loraName ) - if (!ext.prompt.trim()) { + if (!composePromptParts(wrappers.promptPre, ext.prompt, wrappers.promptPost)) { throw new Error(`Shot ${i + 2} needs a prompt`) } @@ -385,7 +387,10 @@ export async function continueQueuedExtensions( throw new Error('Could not extract the last frame for the next extension: the frame file was empty') } job.library.extendPart1Path = part1Path - job.library.prompt = ext.prompt + job.library.prompt = composePromptParts(wrappers.promptPre, ext.prompt, wrappers.promptPost) + job.library.promptMid = ext.prompt + job.library.promptPre = wrappers.promptPre || undefined + job.library.promptPost = wrappers.promptPost || undefined if (live.permanenceRefs?.length) { const refs = [...(job.library.shotPermanenceRefs || [])] refs[i + 1] = live.permanenceRefs diff --git a/utils/promptParts.ts b/utils/promptParts.ts index 82093a0..97cee14 100644 --- a/utils/promptParts.ts +++ b/utils/promptParts.ts @@ -36,6 +36,46 @@ export function restorePromptParts(saved: string, stored?: Partial } } +export type PromptWrapperFields = { + promptPre?: string + promptPost?: string +} + +function hasWrapperKey(item: object | null | undefined, key: 'promptPre' | 'promptPost') { + return Boolean(item && Object.prototype.hasOwnProperty.call(item, key)) +} + +/** First source that actually sent the key wins, including empty string. Omitted keys fall through. */ +export function resolvePromptWrappers( + ...sources: Array +) { + let promptPre = '' + let promptPost = '' + let sawPre = false + let sawPost = false + for (const source of sources) { + if (!source) continue + if (!sawPre && hasWrapperKey(source, 'promptPre') && source.promptPre != null) { + promptPre = sanitizePromptPart(source.promptPre) + sawPre = true + } + if (!sawPost && hasWrapperKey(source, 'promptPost') && source.promptPost != null) { + promptPost = sanitizePromptPart(source.promptPost) + sawPost = true + } + if (sawPre && sawPost) break + } + return { promptPre, promptPost } +} + +export function persistPromptWrappers(item?: PromptWrapperFields | null): PromptWrapperFields { + if (!item) return {} + const next: PromptWrapperFields = {} + if (hasWrapperKey(item, 'promptPre') && item.promptPre != null) next.promptPre = sanitizePromptPart(item.promptPre) + if (hasWrapperKey(item, 'promptPost') && item.promptPost != null) next.promptPost = sanitizePromptPart(item.promptPost) + return next +} + export function formatPromptParts( stored?: { prompt?: string; promptPre?: string; promptPost?: string } | null, midOverride?: string