134 lines
6.1 KiB
TypeScript
134 lines
6.1 KiB
TypeScript
import { normalizeImageIterations } from '~/utils/imageIterations'
|
||
import { assertImageV2LoraStack, parsePostedLoraStack } from '~/server/utils/loras'
|
||
import { persistLoraFields, normalizeLoraStack } from '~/utils/loras'
|
||
import { listStudioJobs, patchStudioJob, type StudioJobPayload } from '~/server/utils/studioQueue'
|
||
|
||
const EDITABLE = new Set(['waiting', 'held', 'error'])
|
||
|
||
function text(raw: unknown, max = 8000) {
|
||
return String(raw ?? '').replace(/\r\n/g, '\n').slice(0, max)
|
||
}
|
||
|
||
function optionalNumber(raw: unknown) {
|
||
if (raw == null || raw === '') return undefined
|
||
const value = Number(raw)
|
||
return Number.isFinite(value) ? value : undefined
|
||
}
|
||
|
||
function optionalBool(raw: unknown) {
|
||
if (raw === true || raw === 'true' || raw === 1 || raw === '1') return true
|
||
if (raw === false || raw === 'false' || raw === 0 || raw === '0') return false
|
||
return undefined
|
||
}
|
||
|
||
export default defineEventHandler(async (event) => {
|
||
const { owner } = assertLibraryOwner(event)
|
||
const id = String(getRouterParam(event, 'id') || '')
|
||
const current = listStudioJobs(owner).find(item => item.id === id)
|
||
if (!current) {
|
||
throw createError({ statusCode: 404, statusMessage: 'Queued job not found' })
|
||
}
|
||
if (current.payload.upscale) throw createError({ statusCode: 409, statusMessage: 'Cancel this upscale and use the clip dialog to change settings.' })
|
||
if (!EDITABLE.has(current.status)) {
|
||
throw createError({ statusCode: 409, statusMessage: 'That job is already generating. Copy it or save a template.' })
|
||
}
|
||
const body = await readBody<Record<string, unknown>>(event).catch(() => ({}))
|
||
|
||
const job = await patchStudioJob(owner, id, (row) => {
|
||
if (!EDITABLE.has(row.status)) throw createError({ statusCode: 409, statusMessage: 'That job has started generating.' })
|
||
const payload = row.payload
|
||
if (body.name != null) {
|
||
const name = text(body.name, 80).trim()
|
||
payload.name = name
|
||
row.name = name || payload.prompt.slice(0, 80)
|
||
}
|
||
if (body.prompt != null) {
|
||
payload.prompt = text(body.prompt)
|
||
row.prompt = payload.prompt
|
||
}
|
||
if (body.tags != null && row.kind === 'music') {
|
||
payload.prompt = text(body.tags, 2000)
|
||
row.prompt = payload.prompt
|
||
}
|
||
if (body.lyrics != null) payload.lyrics = text(body.lyrics)
|
||
const instrumental = optionalBool(body.instrumental)
|
||
if (instrumental != null && row.kind === 'music') {
|
||
if (instrumental) throw createError({ statusCode: 400, statusMessage: 'YuE2 requires lyrics. Instrumental mode is not supported.' })
|
||
payload.instrumental = false
|
||
} else if (instrumental != null) {
|
||
payload.instrumental = instrumental
|
||
}
|
||
if (body.promptPre != null) payload.promptPre = text(body.promptPre)
|
||
if (body.promptMid != null) payload.promptMid = text(body.promptMid)
|
||
if (body.promptPost != null) payload.promptPost = text(body.promptPost)
|
||
if (body.negative != null) payload.negative = text(body.negative, 2000)
|
||
const duration = optionalNumber(body.duration)
|
||
if (duration != null) {
|
||
if (row.kind === 'music') {
|
||
if (!Number.isInteger(duration) || duration < 15 || duration > 150) throw createError({ statusCode: 400, statusMessage: 'YuE2 target length must be 15–150 seconds.' })
|
||
payload.duration = duration
|
||
payload.musicEngine = 'yue2'
|
||
} else payload.duration = Math.min(120, Math.max(0.5, duration))
|
||
}
|
||
const steps = optionalNumber(body.steps)
|
||
if (steps != null) payload.steps = Math.min(100, Math.max(1, Math.round(steps)))
|
||
const cfg = optionalNumber(body.cfg)
|
||
if (cfg != null) payload.cfg = Math.min(20, Math.max(0, Math.round(cfg * 10) / 10))
|
||
const fps = optionalNumber(body.fps)
|
||
if (fps === 12 || fps === 24 || fps === 30) payload.fps = fps
|
||
const seed = optionalNumber(body.seed)
|
||
if (seed != null) payload.seed = Math.max(0, Math.min(2_147_483_647, Math.floor(seed)))
|
||
const turbo = optionalBool(body.turbo)
|
||
if (turbo != null) payload.turbo = turbo
|
||
const sound = optionalBool(body.sound)
|
||
if (sound != null) payload.sound = sound
|
||
if (body.samplerName != null) payload.samplerName = text(body.samplerName, 40)
|
||
if (body.scheduler != null) payload.scheduler = text(body.scheduler, 40)
|
||
if (Array.isArray(body.loraStack) || body.loraName != null) {
|
||
const fields = persistLoraFields(body.loraStack ?? payload.loraStack ?? body.loraName)
|
||
payload.loraName = fields.loraName
|
||
payload.loraStack = fields.loraStack
|
||
}
|
||
if (Array.isArray(body.passes) && row.kind === 'edit') {
|
||
try {
|
||
payload.passes = normalizeImageIterations(body.passes)
|
||
if (payload.imagePipeline === 'v2') {
|
||
for (const pass of payload.passes) {
|
||
if (pass.loraStack !== undefined) pass.loraStack = assertImageV2LoraStack(parsePostedLoraStack(pass.loraStack, 'image'), payload.engine || 'flux')
|
||
}
|
||
}
|
||
} catch (error) {
|
||
throw createError({ statusCode: 400, statusMessage: error instanceof Error ? error.message : String(error) })
|
||
}
|
||
row.shotCount = 1 + payload.passes.length
|
||
}
|
||
if (Array.isArray(body.extensions) && row.kind !== 'edit') {
|
||
payload.extensions = body.extensions.map((item, index) => {
|
||
const rec = (item && typeof item === 'object' ? item : {}) as {
|
||
prompt?: unknown
|
||
duration?: unknown
|
||
loraName?: unknown
|
||
loraStack?: unknown
|
||
}
|
||
const prev = payload.extensions?.[index]
|
||
const fields = persistLoraFields(rec.loraStack || rec.loraName || prev?.loraStack || prev?.loraName)
|
||
const shotDuration = optionalNumber(rec.duration)
|
||
return {
|
||
...prev,
|
||
prompt: text(rec.prompt),
|
||
duration: shotDuration != null ? Math.min(120, Math.max(0.5, shotDuration)) : (prev?.duration || payload.duration),
|
||
loraName: fields.loraName,
|
||
loraStack: fields.loraStack
|
||
}
|
||
}) as StudioJobPayload['extensions']
|
||
row.shotCount = 1 + payload.extensions.length
|
||
}
|
||
if (normalizeLoraStack(payload.loraStack).length && !payload.loraName) {
|
||
payload.loraName = payload.loraStack?.[0]?.name
|
||
}
|
||
if (body.prompt != null && !row.name) row.name = payload.prompt.slice(0, 80)
|
||
})
|
||
|
||
return job
|
||
})
|