Files
aigen/server/utils/loras.ts
T

582 lines
20 KiB
TypeScript

import { comfyConfigured, comfyFetch, getComfyHost } from '~/server/utils/comfy'
import { getBeastImageHost, imageComfyFetch, sameImageHost } from '~/server/utils/imageComfy'
import {
XAIGEN_LORA_MESSAGE,
MINIMAX_TURBO_LORA,
filterLorasForStudio,
imageLoraEngineMismatchMessage,
imageLoraEngineOf,
isXaigenOnlyLora,
loraIdentityKey,
normalizeLoraName,
normalizeLoraStack,
persistLoraFields,
resolveComfyLoraName,
type LoraKind,
type LoraStackItem
} from '~/utils/loras'
import { isXaigenStudio, LTX_DISTILLED_LORA } from '~/utils/videoModels'
type WorkflowNode = { class_type: string; inputs: Record<string, unknown>; _meta?: { title?: string } }
type WorkflowGraph = Record<string, WorkflowNode>
const LORA_LOADER = 'LoraLoader'
const LORA_MODEL_ONLY = 'LoraLoaderModelOnly'
const OBJECT_INFO_NODES = [
LORA_LOADER,
LORA_MODEL_ONLY,
'Power Lora Loader (rgthree)',
'Lora Loader Stack (rgthree)',
'LoraLoaderModelOnly [10]',
'WanVideoLoraSelect',
'LoraLoaderStacked'
]
const MODEL_FEED_CLASSES = new Set(['BasicGuider', 'CFGGuider', 'KSampler', 'KSamplerAdvanced'])
const CACHE_MS = 20_000
const SYSTEM_LORA_PREFERRED = [MINIMAX_TURBO_LORA, LTX_DISTILLED_LORA]
type LoraCache = {
at: number
image: string[]
video: string[]
}
let cache: LoraCache | null = null
let inflight: Promise<LoraCache> | null = null
let inflightFresh = false
function uniqueNames(values: unknown[]) {
const seen = new Set<string>()
const out: string[] = []
for (const value of values) {
const name = String(value || '').trim()
if (!name || seen.has(name)) continue
seen.add(name)
out.push(name)
}
return out
}
function comboFromSpec(spec: unknown): string[] {
if (!Array.isArray(spec)) return []
const first = spec[0]
if (Array.isArray(first)) return first.filter((item): item is string => typeof item === 'string' && item.trim().length > 0)
if (typeof first === 'string' && first.trim()) return spec.filter((item): item is string => typeof item === 'string' && item.trim().length > 0)
return []
}
function namesFromObjectInfoNode(info: unknown) {
if (!info || typeof info !== 'object') return [] as string[]
const input = (info as { input?: { required?: Record<string, unknown>; optional?: Record<string, unknown> } }).input
const buckets = [input?.required, input?.optional]
const names: string[] = []
for (const bucket of buckets) {
if (!bucket) continue
for (const [key, spec] of Object.entries(bucket)) {
if (!/lora/i.test(key)) continue
names.push(...comboFromSpec(spec))
}
}
return uniqueNames(names)
}
async function fetchJson(path: string, via: 'video' | 'image') {
const res = via === 'image'
? await imageComfyFetch(path, { signal: AbortSignal.timeout(8000) })
: await comfyFetch(path, { signal: AbortSignal.timeout(8000) })
if (!res.ok) return null
return res.json().catch(() => null)
}
async function fetchModelsLoras(via: 'video' | 'image') {
const payload = await fetchJson('/models/loras', via)
if (Array.isArray(payload)) return uniqueNames(payload)
if (payload && typeof payload === 'object' && Array.isArray((payload as { loras?: unknown[] }).loras)) {
return uniqueNames((payload as { loras: unknown[] }).loras)
}
return [] as string[]
}
async function fetchObjectInfoLoras(via: 'video' | 'image') {
const image: string[] = []
const video: string[] = []
let any = false
for (const node of OBJECT_INFO_NODES) {
const info = await fetchJson(`/object_info/${encodeURIComponent(node)}`, via)
if (!info || typeof info !== 'object') continue
const record = info as Record<string, unknown>
const body = record[node] || (Object.keys(record).length === 1 ? record[Object.keys(record)[0] as string] : record)
const names = namesFromObjectInfoNode(body)
if (!names.length) continue
any = true
if (node === LORA_LOADER) image.push(...names)
else video.push(...names)
if (node !== LORA_LOADER && node !== LORA_MODEL_ONLY) {
image.push(...names)
video.push(...names)
}
}
if (any) {
return { image: uniqueNames(image), video: uniqueNames(video) }
}
const all = await fetchJson('/object_info', via)
if (!all || typeof all !== 'object') return { image: [] as string[], video: [] as string[] }
const imageAll: string[] = []
const videoAll: string[] = []
for (const [classType, info] of Object.entries(all as Record<string, unknown>)) {
if (!/lora/i.test(classType)) continue
const names = namesFromObjectInfoNode(info)
if (classType === LORA_LOADER) imageAll.push(...names)
else videoAll.push(...names)
}
return { image: uniqueNames(imageAll), video: uniqueNames(videoAll) }
}
async function discoverFromHost(via: 'video' | 'image') {
const models = await fetchModelsLoras(via).catch(() => [] as string[])
const fromInfo = await fetchObjectInfoLoras(via).catch(() => ({ image: [] as string[], video: [] as string[] }))
const image = uniqueNames([...fromInfo.image, ...models])
const video = uniqueNames([...fromInfo.video, ...models])
if (!image.length && video.length) return { image: video, video }
if (!video.length && image.length) return { image, video: image }
return { image, video }
}
async function loadLoraCache(options: { fresh?: boolean } = {}): Promise<LoraCache> {
const fresh = options.fresh === true
const now = Date.now()
if (!fresh && cache && now - cache.at < CACHE_MS) return cache
if (inflight && (!fresh || inflightFresh)) return inflight
inflightFresh = fresh
inflight = (async () => {
const videoHost = comfyConfigured() ? getComfyHost() : ''
const imageHost = getBeastImageHost()
const same = Boolean(videoHost && imageHost && sameImageHost(videoHost, imageHost))
const video = videoHost
? await discoverFromHost('video').catch(() => ({ image: [] as string[], video: [] as string[] }))
: { image: [] as string[], video: [] as string[] }
const image = imageHost && !same
? await discoverFromHost('image').catch(() => ({ image: [] as string[], video: [] as string[] }))
: video
const next: LoraCache = {
at: Date.now(),
image: uniqueNames([...image.image, ...video.image]),
video: uniqueNames([...video.video, ...image.video])
}
// Comfy asleep / unreachable returns empty. Keep the last good list so the picker
// does not vanish — but do NOT refresh `at`, or a blip freezes new folder drops out.
if (!next.image.length && !next.video.length && cache && (cache.image.length || cache.video.length)) {
return cache
}
// Prefer the richer list when a partial reply would shrink a known catalog.
if (
cache
&& (next.image.length + next.video.length) < (cache.image.length + cache.video.length)
&& next.image.every(name => cache!.image.includes(name) || cache!.video.includes(name))
&& next.video.every(name => cache!.image.includes(name) || cache!.video.includes(name))
) {
const merged: LoraCache = {
at: Date.now(),
image: uniqueNames([...cache.image, ...next.image]),
video: uniqueNames([...cache.video, ...next.video])
}
cache = merged
return merged
}
cache = next
return next
})().finally(() => {
inflight = null
inflightFresh = false
})
return inflight
}
export async function listStudioLoras(options: { fresh?: boolean } = {}) {
const xaigen = isXaigenStudio()
try {
const listed = await loadLoraCache(options)
return {
at: listed.at,
image: filterLorasForStudio(listed.image, xaigen),
video: filterLorasForStudio(listed.video, xaigen)
}
} catch {
if (cache) {
return {
at: cache.at,
image: filterLorasForStudio(cache.image, xaigen),
video: filterLorasForStudio(cache.video, xaigen)
}
}
return { at: 0, image: [] as string[], video: [] as string[] }
}
}
/** Unfiltered Comfy filenames for graph system LoRAs (MiniMax turbo, LTX distilled). Reuses the listing cache. */
export function cachedComfyLoraNames(kind?: LoraKind) {
if (!cache) return [] as string[]
if (kind === 'image') return uniqueNames([...cache.image, ...cache.video])
if (kind === 'video') return uniqueNames([...cache.video, ...cache.image])
return uniqueNames([...cache.image, ...cache.video])
}
export async function ensureComfyLoraNames(kind?: LoraKind) {
try {
await loadLoraCache()
} catch {
/* keep whatever was cached */
}
return cachedComfyLoraNames(kind)
}
export function resolveGraphLoraNames(graph: WorkflowGraph, kind?: LoraKind) {
const names = cachedComfyLoraNames(kind)
for (const node of Object.values(graph)) {
if (node.class_type !== LORA_LOADER && node.class_type !== LORA_MODEL_ONLY) continue
const current = String(node.inputs.lora_name || '').trim()
if (!current) continue
const preferred = SYSTEM_LORA_PREFERRED.find(item => loraIdentityKey(item) === loraIdentityKey(current)) || current
node.inputs.lora_name = names.length ? resolveComfyLoraName(preferred, names) : preferred
}
}
function allowedLoraNames(kind?: LoraKind) {
return filterLorasForStudio(cachedComfyLoraNames(kind), isXaigenStudio())
}
function resolveUserLoraName(name: string, kind?: LoraKind) {
const allowed = allowedLoraNames(kind)
if (!allowed.length) return name
return resolveComfyLoraName(name, allowed)
}
export function assertLoraAllowed(raw: unknown, kind: LoraKind) {
const name = normalizeLoraName(raw)
if (!name) return ''
if (isXaigenOnlyLora(name) && !isXaigenStudio()) {
throw createError({ statusCode: 400, statusMessage: XAIGEN_LORA_MESSAGE })
}
const allowed = allowedLoraNames(kind)
if (allowed.length) {
const identity = loraIdentityKey(name)
const listed = allowed.some((item) => {
const path = item.replace(/\\/g, '/')
return path === name || path.toLowerCase() === name.toLowerCase() || loraIdentityKey(item) === identity
})
if (!listed) {
throw createError({ statusCode: 400, statusMessage: `Unknown ${kind} LoRA` })
}
return resolveUserLoraName(name, kind)
}
return name
}
export function parsePostedLora(raw: unknown, kind: LoraKind) {
return parsePostedLoraStack(raw, kind)[0]?.name || ''
}
export function parsePostedLoraStack(raw: unknown, kind: LoraKind): LoraStackItem[] {
const out: LoraStackItem[] = []
for (const item of normalizeLoraStack(raw)) {
const name = assertLoraAllowed(item.name, kind)
if (!name) continue
out.push({ ...item, name })
}
return out
}
export function parseShotLoras(raw: unknown, shotCount: number, kind: LoraKind = 'video') {
return parseShotLoraStacks(raw, shotCount, kind).map(stack => stack[0]?.name || '')
}
export function parseShotLoraStacks(raw: unknown, shotCount: number, kind: LoraKind = 'video'): LoraStackItem[][] {
const empty = Array.from({ length: shotCount }, () => [] as LoraStackItem[])
if (!raw) return empty
let parsed: unknown = raw
if (typeof raw === 'string') {
try {
parsed = JSON.parse(raw)
} catch {
return empty
}
}
if (!Array.isArray(parsed)) return empty
return Array.from({ length: shotCount }, (_, index) => parsePostedLoraStack(parsed[index], kind))
}
export { persistLoraFields }
function linkSource(value: unknown): string | null {
return Array.isArray(value) && typeof value[0] === 'string' ? value[0] : null
}
function alreadyHasLora(graph: WorkflowGraph, name: string) {
const wanted = loraIdentityKey(name)
if (!wanted) return false
return Object.values(graph).some((node) => {
if (node.class_type !== LORA_LOADER && node.class_type !== LORA_MODEL_ONLY) return false
return loraIdentityKey(String(node.inputs.lora_name || '')) === wanted
})
}
function injectAfter(
graph: WorkflowGraph,
sourceId: string,
nodeId: string,
node: WorkflowNode
) {
if (!graph[sourceId] || graph[nodeId]) return
graph[nodeId] = node
for (const [id, other] of Object.entries(graph)) {
if (id === nodeId) continue
for (const [key, value] of Object.entries(other.inputs)) {
const src = linkSource(value)
if (src === sourceId && Array.isArray(value)) {
other.inputs[key] = [nodeId, value[1]]
}
}
}
}
function findModelFeed(graph: WorkflowGraph) {
for (const node of Object.values(graph)) {
if (!MODEL_FEED_CLASSES.has(node.class_type)) continue
const source = linkSource(node.inputs.model)
if (source && graph[source]) return source
}
return ''
}
function bypassLoraNode(graph: WorkflowGraph, id: string) {
const node = graph[id]
if (!node) return
const model = Array.isArray(node.inputs.model) ? node.inputs.model : null
const clip = Array.isArray(node.inputs.clip) ? node.inputs.clip : null
for (const [otherId, other] of Object.entries(graph)) {
if (otherId === id) continue
for (const [key, value] of Object.entries(other.inputs)) {
const src = linkSource(value)
if (src !== id || !Array.isArray(value)) continue
const slot = value[1]
if (slot === 0 && model) other.inputs[key] = [model[0], model[1]]
else if (slot === 1 && clip) other.inputs[key] = [clip[0], clip[1]]
}
}
delete graph[id]
}
/** Optional Klein-style LoraLoader: user stack only. Empty stack bypasses model/CLIP around the loader. */
export function applyOptionalLoraLoaders(graph: WorkflowGraph, stack?: unknown, kind: LoraKind = 'image') {
const loaders = Object.entries(graph).filter(([, node]) => node.class_type === LORA_LOADER)
const xaigen = isXaigenStudio()
const items = normalizeLoraStack(stack)
.filter(item => xaigen || !isXaigenOnlyLora(item.name))
.map(item => ({
...item,
name: resolveUserLoraName(item.name, kind)
}))
.filter(item => item.name && (xaigen || !isXaigenOnlyLora(item.name)))
if (!loaders.length) {
applyUserLoraToGraph(graph, items)
return
}
if (!items.length) {
for (const [id] of loaders) bypassLoraNode(graph, id)
return
}
for (const [id] of loaders) {
const first = items[0]
graph[id].inputs.lora_name = first.name
graph[id].inputs.strength_model = first.strengthModel
graph[id].inputs.strength_clip = first.strengthClip
graph[id]._meta = { title: items.length === 1 ? 'User LoRA' : 'User LoRA 1' }
let sourceId = id
for (const [index, item] of items.slice(1).entries()) {
const nodeId = `user:lora:${id}:${index + 1}`
injectAfter(graph, sourceId, nodeId, {
class_type: LORA_LOADER,
inputs: {
lora_name: item.name,
strength_model: item.strengthModel,
strength_clip: item.strengthClip,
model: [sourceId, 0],
clip: [sourceId, 1]
},
_meta: { title: `User LoRA ${index + 2}` }
})
sourceId = nodeId
}
}
}
function linkRef(value: unknown): [string, number] | null {
return Array.isArray(value) && typeof value[0] === 'string' && Number.isFinite(Number(value[1]))
? [value[0], Number(value[1])]
: null
}
export function assertImageV2LoraStack(stack: unknown, engine: 'flux' | 'krea') {
const items = normalizeLoraStack(stack)
const wrong = items.find((item) => {
const tagged = imageLoraEngineOf(item.name)
return tagged != null && tagged !== engine
})
if (wrong) {
throw createError({ statusCode: 400, statusMessage: imageLoraEngineMismatchMessage(wrong.name, engine) })
}
return items
}
function stripExistingLoraLoaders(graph: WorkflowGraph) {
const remaining = new Set(
Object.entries(graph)
.filter(([, node]) => node.class_type === LORA_LOADER)
.map(([id]) => id)
)
while (remaining.size) {
const id = [...remaining].find(candidate => (
![...remaining].some(other => other !== candidate && linkSource(graph[other]?.inputs.model) === candidate)
)) || [...remaining][0]
bypassLoraNode(graph, id)
remaining.delete(id)
}
}
function modelConsumers(graph: WorkflowGraph) {
return Object.values(graph).filter(node => MODEL_FEED_CLASSES.has(node.class_type) && linkRef(node.inputs.model))
}
function clipConsumers(graph: WorkflowGraph) {
return Object.values(graph).filter(node => node.class_type === 'CLIPTextEncode' && linkRef(node.inputs.clip))
}
/**
* Replace any template LoRA loaders with the posted chain, in list order.
* Strength 0 skips that card (off). Empty chain leaves UNET/CLIP wired straight to the sampler.
*/
export function applyImageV2UserLoras(
graph: WorkflowGraph,
stack?: unknown,
engine: 'flux' | 'krea' = 'flux'
) {
stripExistingLoraLoaders(graph)
const xaigen = isXaigenStudio()
const items = normalizeLoraStack(stack)
.filter(item => xaigen || !isXaigenOnlyLora(item.name))
.map(item => ({
...item,
name: resolveUserLoraName(item.name, 'image')
}))
.filter(item => item.name && (xaigen || !isXaigenOnlyLora(item.name)) && (item.strengthModel !== 0 || item.strengthClip !== 0))
const wrong = items.find((item) => {
const tagged = imageLoraEngineOf(item.name)
return tagged != null && tagged !== engine
})
if (wrong) {
throw createError({ statusCode: 400, statusMessage: imageLoraEngineMismatchMessage(wrong.name, engine) })
}
const guiders = modelConsumers(graph)
const encodes = clipConsumers(graph)
const modelFrom = linkRef(guiders[0]?.inputs.model) || (graph['4'] ? ['4', 0] as [string, number] : null)
const clipFrom = linkRef(encodes[0]?.inputs.clip) || (graph['5'] ? ['5', 0] as [string, number] : null)
if (!modelFrom || !clipFrom) {
throw createError({ statusCode: 500, statusMessage: 'v2 graph has no model/CLIP feed for user LoRAs.' })
}
if (!items.length) return [] as LoraStackItem[]
if (!guiders.length || !encodes.length) {
throw createError({
statusCode: 500,
statusMessage: 'Image v2 LoRA chain is not connected to the sampler. Refusing to run without those adapters.'
})
}
let model: [string, number] = modelFrom
let clip: [string, number] = clipFrom
const ids: string[] = []
for (const [index, item] of items.entries()) {
const nodeId = String(70 + index)
if (graph[nodeId]) {
throw createError({ statusCode: 500, statusMessage: `v2 graph already has node ${nodeId}. Refusing to overwrite it with a LoRA.` })
}
graph[nodeId] = {
class_type: LORA_LOADER,
inputs: {
lora_name: item.name,
strength_model: item.strengthModel,
strength_clip: item.strengthClip,
model,
clip
},
_meta: { title: items.length === 1 ? 'LoRA' : `LoRA ${index + 1}` }
}
ids.push(nodeId)
model = [nodeId, 0]
clip = [nodeId, 1]
}
for (const node of guiders) node.inputs.model = model
for (const node of encodes) node.inputs.clip = clip
const hooked = guiders.every(node => linkRef(node.inputs.model)?.[0] === ids[ids.length - 1])
&& encodes.every(node => linkRef(node.inputs.clip)?.[0] === ids[ids.length - 1])
if (!hooked) {
throw createError({
statusCode: 500,
statusMessage: 'Image v2 LoRA chain is not connected to the sampler. Refusing to run without those adapters.'
})
}
return items
}
export function applyUserLoraToGraph(graph: WorkflowGraph, stack?: unknown) {
const xaigen = isXaigenStudio()
const items = normalizeLoraStack(stack)
.filter(item => xaigen || !isXaigenOnlyLora(item.name))
.map(item => ({ ...item, name: resolveUserLoraName(item.name) }))
.filter(item => item.name && (xaigen || !isXaigenOnlyLora(item.name)) && !alreadyHasLora(graph, item.name))
if (!items.length) return
const clipLoaders = Object.entries(graph).filter(([, node]) => node.class_type === LORA_LOADER)
if (clipLoaders.length) {
for (const [id] of clipLoaders) {
let sourceId = id
for (const [index, item] of items.entries()) {
const nodeId = `user:lora:${id}:${index}`
injectAfter(graph, sourceId, nodeId, {
class_type: LORA_LOADER,
inputs: {
lora_name: item.name,
strength_model: item.strengthModel,
strength_clip: item.strengthClip,
model: [sourceId, 0],
clip: [sourceId, 1]
},
_meta: { title: items.length === 1 ? 'User LoRA' : `User LoRA ${index + 1}` }
})
sourceId = nodeId
}
}
return
}
let source = findModelFeed(graph)
if (!source) return
for (const [index, item] of items.entries()) {
const nodeId = index === 0 ? 'user:lora' : `user:lora:${index}`
injectAfter(graph, source, nodeId, {
class_type: LORA_MODEL_ONLY,
inputs: {
lora_name: item.name,
strength_model: item.strengthModel,
model: [source, 0]
},
_meta: { title: items.length === 1 ? 'User LoRA' : `User LoRA ${index + 1}` }
})
source = nodeId
}
}