diff --git a/web/comfy_shared.js b/web/comfy_shared.js index a33d747..edeabfe 100644 --- a/web/comfy_shared.js +++ b/web/comfy_shared.js @@ -155,12 +155,14 @@ export function getWidgetType(config) { } return { type, linkType } } -export const setupDynamicConnections = (nodeType, prefix, inputType) => { +export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => { + infoLogger('Setting up dynamic connections for', nodeType) + const options = opts || {} const onNodeCreated = nodeType.prototype.onNodeCreated - // check if it's a list const inputList = typeof inputType === 'object' + nodeType.prototype.onNodeCreated = function () { - const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined + const r = onNodeCreated ? onNodeCreated.apply(this) : undefined this.addInput(`${prefix}_1`, inputList ? '*' : inputType) return r } @@ -168,64 +170,208 @@ export const setupDynamicConnections = (nodeType, prefix, inputType) => { const onConnectionsChange = nodeType.prototype.onConnectionsChange nodeType.prototype.onConnectionsChange = function ( type, - index, - connected, - link_info, + slotIndex, + isConnected, + link, + ioSlot, ) { + infoLogger(`Connection changed for ${this.type}`, { + node: this, + type, + slotIndex, + isConnected, + link, + ioSlot, + }) + options.link = link + options.ioSlot = ioSlot + const r = onConnectionsChange - ? onConnectionsChange.apply(this, arguments) + ? onConnectionsChange.apply( + this, + type, + slotIndex, + isConnected, + link, + ioSlot, + ) : undefined - dynamic_connection(this, index, connected, `${prefix}_`, inputList) + dynamic_connection( + this, + slotIndex, + isConnected, + `${prefix}_`, + inputType, + options, + ) + return r } } +/** + * cleanup dynamic inputs + * + * @param {import("../../../web/types/litegraph.d.ts").LGraphNode} node - The target node + * @param {bool} connected - Was this event connecting or disconnecting + * @param {string} connectionPrefix - The common prefix of the dynamic inputs + * @param {string|[string]} connectionType - The type of the dynamic connection + * @param {{nameInput?:[string]}} [opts] - extra options + */ + +const clean_dynamic_state = ( + node, + connected, + connectionPrefix, + connectionType, + opts, +) => { + infoLogger('CLEANING', { node, connectionPrefix, connectionType, opts }) + const options = opts || {} + const nameArray = options.nameArray || [] + + const listConnection = typeof connectionType === 'object' + const conType = listConnection ? '*' : connectionType + infoLogger('connected', connected) + + if (connected) { + // Remove inputs and their widget if not linked. + for (let n = 0; n < node.inputs.length; n++) { + const element = node.inputs[n] + if (!element.link) { + if (node.widgets) { + const w = node.widgets.find((w) => w.name === element.name) + if (w) { + w.onRemoved?.() + node.widgets.length = node.widgets.length - 1 + } + } + node.removeInput(n) + } + } + } + // make inputs sequential again + for (let i = 0; i < node.inputs.length; i++) { + let name = `${connectionPrefix}${i + 1}` + + if (nameArray.length > 0) { + name = i < nameArray.length ? nameArray[i] : name + } + + node.inputs[i].label = name + node.inputs[i].name = name + } + // add an extra input + if (node.inputs[node.inputs.length - 1].link !== undefined) { + const nextIndex = node.inputs.length + let name = `${connectionPrefix}${nextIndex + 1}` + if (nameArray.length > 0) { + name = nextIndex < nameArray.length ? nameArray[nextIndex] : name + } + log(`Adding input ${nextIndex + 1} (${name})`) + node.addInput(name, conType) + } +} + +/** + * Main logic around dynamic inputs + * + * @param {import("../../../web/types/litegraph.d.ts").LGraphNode} node - The target node + * @param {number} index - The slot index of the currently changed connection + * @param {bool} connected - Was this event connecting or disconnecting + * @param {string} [connectionPrefix] - The common prefix of the dynamic inputs + * @param {string|[string]} [connectionType] - The type of the dynamic connection + * @param {{nameInput?:[string]}} [opts] - extra options + */ export const dynamic_connection = ( node, index, connected, connectionPrefix = 'input_', - connectionType = 'PSDLAYER', - nameArray = [], + connectionType = '*', + opts = undefined, ) => { + infoLogger('MTB Dynamic Connection', { + node, + node_inputs: node.inputs, + index, + connected, + connectionPrefix, + connectionType, + opts, + }) + const options = opts || {} if (!node.inputs[index].name.startsWith(connectionPrefix)) { return } + const listConnection = typeof connectionType === 'object' - // remove all non connected inputs - if (!connected && node.inputs.length > 1) { - log(`Removing input ${index} (${node.inputs[index].name})`) - if (node.widgets) { - const w = node.widgets.find((w) => w.name === node.inputs[index].name) - if (w) { - w.onRemoved?.() - node.widgets.length = node.widgets.length - 1 + const conType = listConnection ? '*' : connectionType + const nameArray = options.nameArray || [] + + // clean_dynamic_state( + // node, + // connected, + // connectionPrefix, + // connectionType, + // options, + // ) + // + + if (connected) { + // Remove inputs and their widget if not linked. + for (let n = 0; n < node.inputs.length; n++) { + const element = node.inputs[n] + if (!element.link) { + if (node.widgets) { + const w = node.widgets.find((w) => w.name === element.name) + if (w) { + w.onRemoved?.() + node.widgets.length = node.widgets.length - 1 + } + } + node.removeInput(n) } } - node.removeInput(index) + } + // make inputs sequential again + for (let i = 0; i < node.inputs.length; i++) { + let name = `${connectionPrefix}${i + 1}` - // make inputs sequential again - for (let i = 0; i < node.inputs.length; i++) { - const name = - i < nameArray.length ? nameArray[i] : `${connectionPrefix}${i + 1}` - node.inputs[i].label = name - node.inputs[i].name = name + if (nameArray.length > 0) { + name = i < nameArray.length ? nameArray[i] : name } + + node.inputs[i].label = name + node.inputs[i].name = name } // add an extra input - if (node.inputs[node.inputs.length - 1].link != undefined) { - const nextIndex = node.inputs.length - const name = - nextIndex < nameArray.length - ? nameArray[nextIndex] - : `${connectionPrefix}${nextIndex + 1}` - - log(`Adding input ${nextIndex + 1} (${name})`) - - node.addInput(name, listConnection ? '*' : connectionType) + if (node.inputs.length === 0) { + let name = `${connectionPrefix}1` + if (nameArray.length > 0) { + name = nameArray.length[0] + } + log(`Adding input 1 (${name})`) + node.addInput(name, conType) + } else { + if (node.inputs[node.inputs.length - 1].link !== undefined) { + const nextIndex = node.inputs.length + let name = `${connectionPrefix}${nextIndex + 1}` + if (nameArray.length > 0) { + name = nextIndex < nameArray.length ? nameArray[nextIndex] : name + } + log(`Adding input ${nextIndex + 1} (${name})`) + node.addInput(name, conType) + } } } +/** + * Calculate total height of DOM element child + * + * @param {HTMLElement} parentElement - The target dom element + * @returns {number} the total height + */ export function calculateTotalChildrenHeight(parentElement) { let totalHeight = 0 @@ -233,11 +379,11 @@ export function calculateTotalChildrenHeight(parentElement) { const style = window.getComputedStyle(child) // Get height as an integer (without 'px') - const height = parseInt(style.height, 10) + const height = Number.parseInt(style.height, 10) // Get vertical margin as integers - const marginTop = parseInt(style.marginTop, 10) - const marginBottom = parseInt(style.marginBottom, 10) + const marginTop = Number.parseInt(style.marginTop, 10) + const marginBottom = Number.parseInt(style.marginBottom, 10) // Sum up height and vertical margins totalHeight += height + marginTop + marginBottom @@ -252,9 +398,9 @@ export function calculateTotalChildrenHeight(parentElement) { */ export function addMenuHandler(nodeType, cb) { const getOpts = nodeType.prototype.getExtraMenuOptions - nodeType.prototype.getExtraMenuOptions = function () { - const r = getOpts.apply(this, arguments) - cb.apply(this, arguments) + nodeType.prototype.getExtraMenuOptions = function (node, options) { + const r = getOpts.apply(this, [node, options]) + cb.apply(this, [node, options]) return r } } @@ -280,11 +426,16 @@ export function hideWidget(node, widget, suffix = '') { // Hide any linked widgets, e.g. seed+seedControl if (widget.linkedWidgets) { for (const w of widget.linkedWidgets) { - hideWidget(node, w, ':' + widget.name) + hideWidget(node, w, `:${widget.name}`) } } } +/** + * Show widget + * + * @param {import("../../../web/types/litegraph.d.ts").IWidget} widget - target widget + */ export function showWidget(widget) { widget.type = widget.origType widget.computeSize = widget.origComputeSize @@ -352,7 +503,7 @@ export function hideWidgetForGood(node, widget, suffix = '') { // Hide any linked widgets, e.g. seed+seedControl if (widget.linkedWidgets) { for (const w of widget.linkedWidgets) { - hideWidgetForGood(node, w, ':' + widget.name) + hideWidgetForGood(node, w, `:${widget.name}`) } } } @@ -370,7 +521,7 @@ export function fixWidgets(node) { // continue // } const w = node.widgets.find((w) => w.name === matching_widget.name) - if (w && w.type != CONVERTED_TYPE) { + if (w && w.type !== CONVERTED_TYPE) { log(w) log(`hidding ${w.name}(${w.type}) from ${node.type}`) log(node) @@ -386,15 +537,15 @@ export function fixWidgets(node) { } } export function inner_value_change(widget, value, event = undefined) { - if (widget.type == 'number' || widget.type == 'BBOX') { - value = Number(value) - } else if (widget.type == 'BOOL') { - value = Boolean(value) + let corrected_value = value + if (widget.type === 'number' || widget.type === 'BBOX') { + corrected_value = Number(value) + } else if (widget.type === 'BOOL') { + corrected_value = Boolean(value) } - widget.value = value + widget.value = corrected_value if ( - widget.options && - widget.options.property && + widget.options?.property && node.properties[widget.options.property] !== undefined ) { node.setProperty(widget.options.property, value) @@ -412,13 +563,12 @@ export function isColorBright(rgb, threshold = 240) { function getBrightness(rgbObj) { return Math.round( - (parseInt(rgbObj[0]) * 299 + - parseInt(rgbObj[1]) * 587 + - parseInt(rgbObj[2]) * 114) / + (Number.parseInt(rgbObj[0]) * 299 + + Number.parseInt(rgbObj[1]) * 587 + + Number.parseInt(rgbObj[2]) * 114) / 1000, ) } - //- HTML / CSS UTILS export const loadScript = ( FILE_URL, @@ -439,14 +589,14 @@ export const loadScript = ( scriptEle.async = async scriptEle.src = FILE_URL - scriptEle.addEventListener('load', (ev) => { + scriptEle.addEventListener('load', (_ev) => { resolve({ status: true }) }) - scriptEle.addEventListener('error', (ev) => { + scriptEle.addEventListener('error', (_ev) => { reject({ status: false, - message: `Failed to load the script ${FILE_URL}`, + message: `Failed to load the script ${FILE_URL}`, }) }) @@ -493,7 +643,7 @@ export function defineClass(className, classStyles) { /** Prefixes the node title with '[DEPRECATED]' and log the deprecation reason to the console.*/ export const addDeprecation = (nodeType, reason) => { const title = nodeType.title - nodeType.title = '[DEPRECATED] ' + title + nodeType.title = `[DEPRECATED] ${title}` // console.log(nodeType) const styles = { diff --git a/web/constant.js b/web/constant.js index 900082e..41facaf 100644 --- a/web/constant.js +++ b/web/constant.js @@ -1,62 +1,225 @@ import { app } from '../../scripts/app.js' import * as shared from './comfy_shared.js' +import { infoLogger } from './comfy_shared.js' import { MtbWidgets } from './mtb_widgets.js' +import { ComfyWidgets } from '../../scripts/widgets.js' +import * as mtb_widgets from './mtb_widgets.js' -export class Constant extends LiteGraph.LGraphNode { - constructor() { - super() - this.uuid = shared.makeUUID() - this.collapsable = true +/** + * @typedef {'number'|'string'|'vector2'|'vector3'|'vector4'|'color'} ConstantType + * @typedef {import ("../../../web/types/litegraph.d.ts").LGraphNode} Node + * @typedef {{x:number,y:number,z?:number,w?:number}} VectorValue + * @typedef {} + * + */ - // this avoid serializing the node when converting to prompt - this.isVirtualNode = true - - this.shape = LiteGraph.BOX_SHAPE - this.serialize_widgets = true - - // Properties - this.addProperty('type', 'number') - this.addProperty('value', 0) - - // Inputs and outputs - this.addOutput('Output', '*') - - // Widget for selecting the type - this.addWidget( - 'combo', - 'Type', - this.properties.type, - (value) => { - this.properties.type = value - this.updateWidgets() - this.updateOutputType() - }, - { - values: ['number', 'string', 'vector2', 'vector3', 'vector4', 'color'], - }, - ) - - // Initialize the node - this.updateWidgets() - this.updateOutputType() +/** + * @param {number} size - The number of axis of the vector (2,3 or 4) + * @param {number} val - The default scalar value to fill the vector with + * @returns {VectorValue} vector + * */ +const initVector = (size, val = 0.0) => { + const res = {} + for (let i = 0; i < size; i++) { + const axis = mtb_widgets.VECTOR_AXIS[i] + res[axis] = val } + return res +} + +/** + * + * @extends {Node} + * @classdesc Wrapper for the python node + */ +export class ConstantJs { + constructor(python_node) { + // this.uuid = shared.makeUUID() + const wrapper = this + + python_node.shape = LiteGraph.BOX_SHAPE + python_node.serialize_widgets = true + + const onNodeCreated = python_node.prototype.onNodeCreated + python_node.prototype.onNodeCreated = function () { + const r = onNodeCreated ? onNodeCreated.apply(this) : undefined + + this.addProperty('type', 'number') + this.addProperty('value', 0) + + this.removeInput(0) + this.removeOutput(0) + + this.addOutput('Output', '*') + + // bind our wrapper + this.configure = wrapper.configure.bind(this) + // this.applyToGraph = wrapper.applyToGraph.bind(this) + this.updateWidgets = wrapper.updateWidgets.bind(this) + this.convertValue = wrapper.convertValue.bind(this) + // this.updateOutput = wrapper.updateOutput.bind(this) + this.updateOutputType = wrapper.updateOutputType.bind(this) + // this.updateTargetWidgets = wrapper.updateTargetWidgets.bind(this) + + this.addWidget( + 'combo', + 'Type', + this.properties.type, + (value) => { + this.properties.type = value + this.updateWidgets() + this.updateOutputType() + }, + { + values: [ + // 'number', + 'float', + 'int', + 'string', + 'vector2', + 'vector3', + 'vector4', + 'color', + ], + }, + ) + this.updateWidgets() + this.updateOutputType() + + for (let n = 0; n < this.inputs.length; n++) { + this.removeInput(n) + } + this.inputs = [] + return r + } + return + } + // NOTE: this is called onPrompt - applyToGraph() { - this.updateTargetWidgets() - } + // applyToGraph() { + // infoLogger('Updating values for backend') + // this.updateTargetWidgets() + // } + // NOTE: deserialization happens here configure(info) { - super.configure(info) + // super.configure(info) + infoLogger('Configure Constant', { info, node: this }) + this.properties.type = info.properties.type this.properties.value = info.properties.value - shared.infoLogger('Configure Constant', { info, node: this }) + this.pos = info.pos + this.order = info.order + this.updateWidgets() this.updateOutputType() } + /** + * Convert the old value type to the new one, falling back to some default + * @param {ConstantType} propType - The target type + */ + convertValue(propType) { + switch (propType) { + case 'color': { + if (typeof this.properties.value !== 'string') { + this.properties.value = '#ffffff' + } else if (this.properties.value[0] !== '#') { + this.properties.value = '#ff0000' + } + break + } + case 'int': { + if (typeof this.properties.value === 'object') { + this.properties.value = Number.parseInt(this.properties.value.x) + } else { + this.properties.value = Number.parseInt(this.properties.value) || 0 + } + break + } + case 'float': { + if (typeof this.properties.value === 'object') { + this.properties.value = Number.parseFloat(this.properties.value.x) + } else { + this.properties.value = + Number.parseFloat(this.properties.value) || 0.0 + } + break + } + case 'string': { + if (typeof this.properties.value !== 'string') { + this.properties.value = JSON.stringify(this.properties.value) + } + break + } + case 'vector2': + case 'vector3': + case 'vector4': { + const numInputs = Number.parseInt(propType.charAt(6)) + if (!this.properties.value) { + this.properties.value = initVector(numInputs) // Array.from({ length: numInputs }, () => 0.0) + } else if (typeof this.properties.value === 'string') { + try { + const parsed = JSON.parse(this.properties.value) + const newVec = {} + for ( + let i = 0; + i < Object.keys(mtb_widgets.VECTOR_AXIS).length; + i++ + ) { + const axis = mtb_widgets.VECTOR_AXIS[i] + if (Object.keys(parsed).includes(axis)) { + newVec[axis] = parsed[axis] + } + } + this.properties.value = newVec + } catch (e) { + shared.errorLogger(e) + infoLogger( + `Couldn't parse string to vec (${this.properties.value})`, + ) + this.properties.value = initVector(numInputs) + } + } else if (typeof this.properties.value === 'number') { + const newVec = initVector(numInputs) + newVec.x = Number.parseFloat(this.properties.value) + this.properties.value = newVec + } + + if ( + typeof this.properties.value === 'object' && + Object.keys(this.properties.value).length !== numInputs + ) { + const current = Object.keys(this.properties.value) + if (current.length < numInputs) { + infoLogger('current value smaller than target, adjusting') + for (let index = current.length; index < numInputs; index++) { + this.properties.value[mtb_widgets.VECTOR_AXIS[index]] = 0.0 + } + } else { + infoLogger('current value greater than target, adjusting') + const newVal = {} + for (let index = 0; index < numInputs; index++) { + newVal[mtb_widgets.VECTOR_AXIS[index]] = + this.properties.value[mtb_widgets.VECTOR_AXIS[index]] + } + this.properties.value = newVal + } + } + break + } + default: + break + } + } + + /** + * Remove all widgets but the comboBox for selecting the type + * then recreate the appropriate widget from scratch + */ updateWidgets() { - // Remove existing widgets + // NOTE: Remove existing widgets for (let i = 1; i < this.widgets.length; i++) { const element = this.widgets[i] if (element.onRemove) { @@ -66,28 +229,113 @@ export class Constant extends LiteGraph.LGraphNode { } this.widgets.splice(1) + this.widgets[0].value = this.properties.type + + this.convertValue(this.properties.type) + switch (this.properties.type) { case 'color': { - if (typeof this.properties.value !== 'string') { - this.properties.value = '#ffffff' - } const col_widget = this.addCustomWidget( - MtbWidgets.COLOR('Value', this.properties.value || '#ff0000'), + MtbWidgets.COLOR('Value', this.properties.value), ) col_widget.callback = (col) => { this.properties.value = col - this.updateOutput() + // this.updateOutput() } break } - case 'number': + case 'int': { + const f_widget = this.addCustomWidget( + ComfyWidgets.INT( + this, + 'Value', + [ + '', + { + default: this.properties.value, + callback: (val) => console.log('VALUE', val), + }, + ], + app, + ), + ) + + f_widget.widget.callback = (val) => { + this.properties.value = val + } + + break + } + case 'float': { + this.addWidget('number', 'Value', this.properties.value, (val) => { + this.properties.value = val + }) + break + } + case 'string': { + mtb_widgets.addMultilineWidget( + this, + 'Value', + { + defaultVal: this.properties.value, + }, + (v) => { + this.properties.value = v + // this.updateOutput() + }, + ) + break + } + case 'vector2': + case 'vector3': + case 'vector4': { + const numInputs = Number.parseInt(this.properties.type.charAt(6)) + const node = this + const v_widget = mtb_widgets.addVectorWidget( + this, + 'Value', + this.properties.value, // value + numInputs, // vector_size + function (v) { + node.properties.value = v + // this.updateOutput() + }, + ) + break + } + + // NOTE: this is not reached anymore, kept for reference + case 'number': { if (typeof this.properties.value !== 'number') { this.properties.value = 0.0 } - this.addWidget('number', 'Value', this.properties.value, (value) => { - this.properties.value = value - this.updateOutput() - }) + const n_widget = this.addWidget( + 'number', + 'Value', + this.properties.force_int + ? Number.parseInt(this.properties.value) + : this.properties.value, + (value) => { + this.properties.value = this.properties.force_int + ? Number.parseInt(value) + : value + // this.updateOutput() + }, + ) + //override the callback + const origCallback = n_widget.callback + const node = this + n_widget.callback = function (val) { + const r = origCallback ? origCallback.apply(this, [val]) : undefined + if (node.properties.force_int) { + // TODO: rework this, a it makes it harder to manipulate + this.value = Number.parseInt(this.value) + node.properties.value = Number.parseInt(this.value) + } + infoLogger('NEW NUMBER', this.value) + return r + } + this.addWidget( 'toggle', 'Convert to Integer', @@ -98,51 +346,6 @@ export class Constant extends LiteGraph.LGraphNode { }, ) break - case 'string': { - if (typeof this.properties.value !== 'string') { - this.properties.value = `${this.properties.value}` - } - shared.addMultilineWidget( - this, - 'Value', - { - defaultVal: this.properties.value, - }, - (v) => { - this.properties.value = v - this.updateOutput() - }, - ) - break - } - case 'vector2': - case 'vector3': - case 'vector4': { - const numInputs = Number.parseInt(this.properties.type.charAt(6)) - - if (['string', 'number'].includes(typeof this.properties.value)) { - this.properties.value = Array.from({ length: numInputs }, () => 0.0) - } else if (this.properties.value.length !== numInputs) { - if (this.properties.value.length > numInputs) { - this.properties.value = this.properties.value.slice(0, numInputs) - } else { - this.properties.value = this.properties.value.concat( - new Array(numInputs - this.properties.value.length).fill(0.0), - ) - } - } - for (let i = 0; i < numInputs; i++) { - this.addWidget( - 'number', - `Value ${i + 1}`, - this.properties.value[i] || 0, - (value) => { - this.properties.value[i] = value - this.updateOutput() - }, - ) - } - break } default: break @@ -154,20 +357,28 @@ export class Constant extends LiteGraph.LGraphNode { this.updateTargetWidgets([link.id]) } } + updateOutputType() { - const cur_type = this.outputs[0].type + infoLogger('Updating output type') const rm_if_mismatch = (type) => { - if (cur_type !== type) { + if (this.outputs[0].type !== type) { for (let i = 0; i < this.outputs.length; i++) { this.removeOutput(i) } this.addOutput('output', type) + // this.setOutputDataType(0, type) } } switch (this.properties.type) { case 'color': rm_if_mismatch('COLOR') break + case 'float': + rm_if_mismatch('FLOAT') + break + case 'int': + rm_if_mismatch('INT') + break case 'number': if (this.properties.force_int) { rm_if_mismatch('INT') @@ -178,6 +389,11 @@ export class Constant extends LiteGraph.LGraphNode { case 'string': rm_if_mismatch('STRING') break + // case 'vector2': + // case 'vector3': + // case 'vector4': + // rm_if_mismatch('FLOAT') + // break case 'vector2': rm_if_mismatch('VECTOR2') break @@ -190,7 +406,7 @@ export class Constant extends LiteGraph.LGraphNode { default: break } - this.updateOutput() + // this.updateOutput() } /** @@ -198,6 +414,7 @@ export class Constant extends LiteGraph.LGraphNode { * since Constant is a virtual node. */ updateTargetWidgets(u_links) { + infoLogger('Updating target widgets') if (!app.graph.links) return const links = u_links || this.outputs[0].links if (!links) return @@ -210,12 +427,16 @@ export class Constant extends LiteGraph.LGraphNode { const tgt_widget = tgt_node.widgets.filter( (w) => w.name === tgt_input.name, ) - if (!tgt_widget) return + // infoLogger('Constant Target Node', tgt_node) + // infoLogger('Constant Target Input', tgt_input) + if (!tgt_widget || tgt_widget.length === 0) return + tgt_widget[0].value = this.properties.value } } updateOutput() { + infoLogger('Updating output value') const value = this.properties.value switch (this.properties.type) { @@ -223,40 +444,54 @@ export class Constant extends LiteGraph.LGraphNode { this.setOutputData(0, value) break case 'number': - this.setOutputData(0, Number.parseFloat(value)) + if (this.properties.force_int) { + this.setOutputData(0, Number.parseInt(value)) + } else { + this.setOutputData(0, Number.parseFloat(value)) + } break case 'string': this.setOutputData(0, value.toString()) break case 'vector2': - if (value.length >= 2) { - this.setOutputData(0, value.slice(0, 2)) - } - break case 'vector3': - if (value.length >= 3) { - this.setOutputData(0, value.slice(0, 3)) - } - break case 'vector4': - if (value.length >= 4) { - this.setOutputData(0, value.slice(0, 4)) - } + this.setOutputData(0, value) break + + // case 'vector2': + // this.setOutputData(0, value.slice(0, 2)) + // break + // case 'vector3': + // this.setOutputData(0, value.slice(0, 3)) + // break + // case 'vector4': + // this.setOutputData(0, value.slice(0, 4)) + // break default: break } + infoLogger('New Value', this.value) + this.updateTargetWidgets() } } +app.registerExtension({ + name: 'mtb.constant', -// app.registerExtension({ -// name: 'mtb.constant', -// registerCustomNodes() { -// LiteGraph.registerNodeType('Constant (mtb)', Constant) -// -// Constant.category = 'mtb/utils' -// Constant.title = 'Constant (mtb)' -// }, -// }) + async beforeRegisterNodeDef(nodeType, nodeData, _app) { + if (nodeData.name === 'Constant (mtb)') { + infoLogger('registering constant') + new ConstantJs(nodeType) + } + }, + // NOTE: old js only registration + // + // registerCustomNodes() { + // LiteGraph.registerNodeType('Constant (mtb)', Constant) + // + // Constant.category = 'mtb/utils' + // Constant.title = 'Constant (mtb)' + // }, +}) diff --git a/web/curve_widget.js b/web/curve_widget.js index 39a4b0b..01c0230 100644 --- a/web/curve_widget.js +++ b/web/curve_widget.js @@ -1,187 +1,214 @@ - import { app } from '../../scripts/app.js' -function B0(t) { return (1 - t) ** 3 / 6; } -function B1(t) { return (3 * t ** 3 - 6 * t ** 2 + 4) / 6; } -function B2(t) { return (-3 * t ** 3 + 3 * t ** 2 + 3 * t + 1) / 6; } -function B3(t) { return t ** 3 / 6; } +function B0(t) { + return (1 - t) ** 3 / 6 +} +function B1(t) { + return (3 * t ** 3 - 6 * t ** 2 + 4) / 6 +} +function B2(t) { + return (-3 * t ** 3 + 3 * t ** 2 + 3 * t + 1) / 6 +} +function B3(t) { + return t ** 3 / 6 +} class CurveWidget { - constructor(inputName, defaultValue) { - this.name = inputName || "Curve"; - this._value = defaultValue || [{ x: 0, y: 0 }, { x: 1, y: 1 }]; - this.type = "FLOAT_CURVE"; - this.selectedPointIndex = null; - this.resize - } + constructor(inputName, defaultValue) { + this.name = inputName || 'Curve' + this._value = defaultValue || [ + { x: 0, y: 0 }, + { x: 1, y: 1 }, + ] + this.type = 'FLOAT_CURVE' + this.selectedPointIndex = null + console.log(this) + this.resize() + } - drawBSpline(ctx, width, height, posY) { - const n = this._value.length - 1; - const numSegments = n - 2; - const numPoints = this._value.length; - if (numPoints < 4) { - this.drawLinear(ctx, width, height, posY); - } else { - for (let j = 0; j <= numSegments; j++) { - for (let t = 0; t <= 1; t += 0.01) { - let pt = this.getBSplinePoint(j, t); - let x = pt.x * width; - let y = posY + height - pt.y * height; + drawBSpline(ctx, width, height, posY) { + const n = this._value.length - 1 + const numSegments = n - 2 + const numPoints = this._value.length + if (numPoints < 4) { + this.drawLinear(ctx, width, height, posY) + } else { + for (let j = 0; j <= numSegments; j++) { + for (let t = 0; t <= 1; t += 0.01) { + let pt = this.getBSplinePoint(j, t) + let x = pt.x * width + let y = posY + height - pt.y * height - if (t === 0) ctx.moveTo(x, y); - else ctx.lineTo(x, y); - } - } - ctx.stroke(); + if (t === 0) ctx.moveTo(x, y) + else ctx.lineTo(x, y) } + } + ctx.stroke() } + } - drawLinear(ctx, width, height, posY) { - for (let i = 0; i < this._value.length - 1; i++) { - let p1 = this._value[i]; - let p2 = this._value[i + 1]; - ctx.moveTo(p1.x * width, posY + height - p1.y * height); - ctx.lineTo(p2.x * width, posY + height - p2.y * height); - } - ctx.stroke(); + drawLinear(ctx, width, height, posY) { + for (let i = 0; i < this._value.length - 1; i++) { + let p1 = this._value[i] + let p2 = this._value[i + 1] + ctx.moveTo(p1.x * width, posY + height - p1.y * height) + ctx.lineTo(p2.x * width, posY + height - p2.y * height) } + ctx.stroke() + } - getBSplinePoint(i, t) { - // Control points for this segment - const p0 = this._value[i]; - const p1 = this._value[i + 1]; - const p2 = this._value[i + 2]; - const p3 = this._value[i + 3]; + getBSplinePoint(i, t) { + // Control points for this segment + const p0 = this._value[i] + const p1 = this._value[i + 1] + const p2 = this._value[i + 2] + const p3 = this._value[i + 3] - const x = B0(t) * p0.x + B1(t) * p1.x + B2(t) * p2.x + B3(t) * p3.x; - const y = B0(t) * p0.y + B1(t) * p1.y + B2(t) * p2.y + B3(t) * p3.y; + const x = B0(t) * p0.x + B1(t) * p1.x + B2(t) * p2.x + B3(t) * p3.x + const y = B0(t) * p0.y + B1(t) * p1.y + B2(t) * p2.y + B3(t) * p3.y - return { x, y }; + return { x, y } + } + + draw(ctx, node, width, posY, height) { + const [cw, ch] = this.computeSize(width) + + ctx.beginPath() + ctx.fillStyle = '#000' + //ctx.fillRect(0, posY, cw, ch); + ctx.strokeStyle = '#fff' + ctx.lineWidth = 2 + + // normalized coordinates -> canvas coordinates + for (let i = 0; i < this._value.length - 1; i++) { + let p1 = this._value[i] + let p2 = this._value[i + 1] + ctx.moveTo(p1.x * cw, posY + ch - p1.y * ch) + ctx.lineTo(p2.x * cw, posY + ch - p2.y * ch) } + ctx.stroke() + // this.drawBSpline(ctx, width, height, posY); - draw(ctx, node, width, posY, height) { - const [cw, ch] = this.computeSize(width) + // points + this._value.forEach((point) => { + ctx.beginPath() + ctx.arc(point.x * cw, posY + ch - point.y * ch, 5, 0, 2 * Math.PI) + ctx.fill() + }) + } - ctx.beginPath(); - ctx.fillStyle = "#000"; - //ctx.fillRect(0, posY, cw, ch); - ctx.strokeStyle = "#fff"; - ctx.lineWidth = 2; + mouse(event, pos, node) { + // console.debug(event.type, pos, node) + let x = pos[0] - node.pos[0] + let y = pos[1] - node.pos[1] + let width = node.size[0] + const height = 300 // TODO: compute + const posY = node.pos[1] - // normalized coordinates -> canvas coordinates - for (let i = 0; i < this._value.length - 1; i++) { - let p1 = this._value[i]; - let p2 = this._value[i + 1]; - ctx.moveTo(p1.x * cw, posY + ch - p1.y * ch); - ctx.lineTo(p2.x * cw, posY + ch - p2.y * ch); - } - ctx.stroke(); - // this.drawBSpline(ctx, width, height, posY); + const localPos = { x: pos[0], y: pos[1] - LiteGraph.NODE_WIDGET_HEIGHT } - // points - this._value.forEach(point => { - ctx.beginPath(); - ctx.arc(point.x * cw, posY + ch - point.y * ch, 5, 0, 2 * Math.PI); - ctx.fill(); - }); + if (event.type === LiteGraph.pointerevents_method + 'down') { + console.debug('Checking if a point was clicked') + const clickedPointIndex = this.detectPoint(localPos, width, height) + if (clickedPointIndex !== null) { + this.selectedPointIndex = clickedPointIndex + } else { + this.addPoint(localPos, width, height) + } + return true + } else if ( + event.type === LiteGraph.pointerevents_method + 'move' && + this.selectedPointIndex !== null + ) { + this.movePoint(this.selectedPointIndex, localPos, width, height) + return true + } else if ( + event.type === LiteGraph.pointerevents_method + 'up' && + this.selectedPointIndex !== null + ) { + this.selectedPointIndex = null + return true } + return false + } + callback(...args) { + //value, that, node, pos, event) { - mouse(event, pos, node) { - // console.debug(event.type, pos, node) - let x = pos[0] - node.pos[0] - let y = pos[1] - node.pos[1] - let width = node.size[0] - const height = 300; // TODO: compute - const posY = node.pos[1]; + console.log(args) + } - const localPos = { x: pos[0], y: pos[1] - LiteGraph.NODE_WIDGET_HEIGHT }; - - if (event.type === LiteGraph.pointerevents_method + "down") { - console.debug("Checking if a point was clicked"); - const clickedPointIndex = this.detectPoint(localPos, width, height); - if (clickedPointIndex !== null) { - this.selectedPointIndex = clickedPointIndex; - } else { - this.addPoint(localPos, width, height); - } - return true; - } else if (event.type === LiteGraph.pointerevents_method + "move" && this.selectedPointIndex !== null) { - this.movePoint(this.selectedPointIndex, localPos, width, height); - return true; - } else if (event.type === LiteGraph.pointerevents_method + "up" && this.selectedPointIndex !== null) { - this.selectedPointIndex = null; - return true; - } - return false; + detectPoint(localPos, width, height) { + const threshold = 20 // TODO: extract + for (let i = 0; i < this._value.length; i++) { + const p = this._value[i] + const px = p.x * width + const py = height - p.y * height + if ( + Math.abs(localPos.x - px) < threshold && + Math.abs(localPos.y - py) < threshold + ) { + return i + } } + return null + } - - detectPoint(localPos, width, height) { - const threshold = 20; // TODO: extract - for (let i = 0; i < this._value.length; i++) { - const p = this._value[i]; - const px = p.x * width; - const py = height - p.y * height; - if (Math.abs(localPos.x - px) < threshold && Math.abs(localPos.y - py) < threshold) { - return i; - } - } - return null; + addPoint(localPos, width, height) { + // add a new point based on click position + const normalizedPoint = { + x: localPos.x / width, + y: 1 - localPos.y / height, } + this._value.push(normalizedPoint) + this._value.sort((a, b) => a.x - b.x) + this.value = JSON.stringify(this._value) + } - addPoint(localPos, width, height) { - // add a new point based on click position - const normalizedPoint = { x: localPos.x / width, y: 1 - localPos.y / height }; - this._value.push(normalizedPoint); - this._value.sort((a, b) => a.x - b.x); - this.value = JSON.stringify(this._value); - } + movePoint(index, localPos, width, height) { + const point = this._value[index] + point.x = Math.max(0, Math.min(1, localPos.x / width)) + point.y = Math.max(0, Math.min(1, 1 - localPos.y / height)) - movePoint(index, localPos, width, height) { - const point = this._value[index]; - point.x = Math.max(0, Math.min(1, localPos.x / width)); - point.y = Math.max(0, Math.min(1, 1 - localPos.y / height)); + this._value[index] = point + this.value = JSON.stringify(this._value) + } - this._value[index] = point; - this.value = JSON.stringify(this._value); - } + computeSize(width) { + return [width, 300] + } - computeSize(width) { - return [width, 300]; - } + configure(data) { + console.log('CONFIGURE CURVES', data) + } - configure(data) { - console.log(data) - } - - value() { - console.debug('Returning value', this._value) - return this._value - } - setValue(value) { - console.debug('Setting value', value) - this._value = value - } + value() { + console.debug('Returning value', this._value) + return this._value + } + setValue(value) { + console.debug('Setting value', value) + this._value = value + } } app.registerExtension({ - name: 'mtb.curves', - getCustomWidgets: function () { + name: 'mtb.curves', + getCustomWidgets: function () { + return { + FLOAT_CURVE: (node, inputName, inputData, app) => { + console.debug('Registering float curve widget') + console.log({ inputData }) + const wid = node.addCustomWidget( + new CurveWidget(inputName, inputData[1]?.default), + ) + + console.log(node) return { - FLOAT_CURVE: (node, inputName, inputData, app) => { - console.debug('Registering float curve widget'); - - return { - widget: node.addCustomWidget( - new CurveWidget(inputName, inputData[1]?.default) - ), - minWidth: 150, - minHeight: 30, - } - }, - - + widget: wid, + minWidth: 150, + minHeight: 30, } - }, - + }, + } + }, }) diff --git a/web/debug.js b/web/debug.js index 0a00387..5a22425 100644 --- a/web/debug.js +++ b/web/debug.js @@ -51,9 +51,11 @@ app.registerExtension({ //- infer type if (link_info) { - const fromNode = this.graph._nodes.find( - (otherNode) => otherNode.id === link_info.origin_id, - ) + // const fromNode = this.graph._nodes.find( + // (otherNode) => otherNode.id === link_info.origin_id, + // ) + const fromNode = app.graph.getNodeById(link_info.origin_id) + if (!fromNode) return const type = fromNode.outputs[link_info.origin_slot].type this.inputs[index].type = type // this.inputs[index].label = type.toLowerCase() @@ -83,6 +85,7 @@ app.registerExtension({ this.widgets.length = 1 } let widgetI = 1 + console.log(message) if (message.text) { for (const txt of message.text) { const w = this.addCustomWidget( diff --git a/web/mtb_widgets.js b/web/mtb_widgets.js index 4ce997a..6f2a67c 100644 --- a/web/mtb_widgets.js +++ b/web/mtb_widgets.js @@ -7,6 +7,11 @@ * */ +/** + * @typedef {import("../../../web/types/litegraph.d.ts").IWidget} IWidget + * @typedef {import("../../../web/types/litegraph.d.ts").IWidget} VectorWidget + */ + // TODO: Use the builtin addDOMWidget everywhere appropriate import { app } from '../../scripts/app.js' @@ -15,7 +20,7 @@ import { api } from '../../scripts/api.js' import parseCss from './extern/parse-css.js' import * as shared from './comfy_shared.js' import { log } from './comfy_shared.js' -import { Constant } from './constant.js' +import { NumberInputWidget } from './numberInput.js' // NOTE: new widget types registered by MTB Widgets const newTypes = [/*'BOOL'n,*/ 'COLOR', 'BBOX'] @@ -55,7 +60,250 @@ const calculateTextDimensions = (ctx, value, width, fontSize = 16) => { return { textHeight, maxLineWidth } } +export function addMultilineWidget(node, name, opts, callback) { + const inputEl = document.createElement('textarea') + inputEl.className = 'comfy-multiline-input' + inputEl.value = opts.defaultVal + inputEl.placeholder = opts.placeholder || name + + const widget = node.addDOMWidget(name, 'textmultiline', inputEl, { + getValue() { + return inputEl.value + }, + setValue(v) { + inputEl.value = v + }, + }) + widget.inputEl = inputEl + + inputEl.addEventListener('input', () => { + callback?.(widget.value) + widget.callback?.(widget.value) + }) + widget.onRemove = () => { + inputEl.remove() + } + + return { minWidth: 400, minHeight: 200, widget } +} + +export const VECTOR_AXIS = { + 0: 'x', + 1: 'y', + 2: 'z', + 3: 'w', +} + +export function addVectorWidgetW( + node, + name, + value, + vector_size, + callback, + app, +) { + // const inputEl = document.createElement('div') + // const vecEl = document.createElement('div') + // + // inputEl.style.background = 'red' + // + // inputEl.className = 'comfy-vector-container' + // vecEl.className = 'comfy-vector-input' + // + // vecEl.style.display = 'flex' + // inputEl.appendChild(vecEl) + const inputs = [] + + for (let i = 0; i < vector_size; i++) { + // const input = document.createElement('input') + // input.type = 'number' + // input.value = value[VECTOR_AXIS[i]] + const input = node.addWidget( + 'number', + `${name}_${VECTOR_AXIS[i]}`, + value[VECTOR_AXIS[i]], + (val) => {}, + ) + + inputs.push(input) + // vecEl.appendChild(input) + } + // + // const widget = node.addDOMWidget(name, 'vector', inputEl, { + // getValue() { + // return JSON.stringify(widget._value) + // }, + // setValue(v) { + // widget._value = v + // }, + // afterResize(node, widget) { + // console.log('After resize', { that: this, node, widget }) + // }, + // }) + // + // console.log('prev callback', widget.callback) + // widget.callback = callback + // widget._value = value + // + // for (let i = 0; i < vector_size; i++) { + // const input = inputs[i] + // input.addEventListener('change', (event) => { + // widget._value[VECTOR_AXIS[i]] = Number.parseFloat(event.target.value) + // widget.callback?.(widget._value) + // node.graph._version++ + // node.setDirtyCanvas(true, true) + // }) + // } + // // document.body.append(inputEl) + // + // widget.inputEl = inputEl + // widget.vecEl = vecEl + // + // inputEl.addEventListener('input', () => { + // widget.callback?.(widget.value) + // }) + // + return { minWidth: 400, minHeight: 200, widget } +} +export function addVectorWidget(node, name, value, vector_size, callback, app) { + const inputEl = document.createElement('div') + const vecEl = document.createElement('div') + + inputEl.className = 'comfy-vector-container' + vecEl.className = 'comfy-vector-input' + vecEl.id = 'vecEl' + + vecEl.style.display = 'flex' + vecEl.style.flexDirection = 'column' + inputEl.appendChild(vecEl) + const inputs = [] + + // + // for (let i = 0; i < vector_size; i++) { + // const input = document.createElement('input') + // input.type = 'number' + // input.value = value[VECTOR_AXIS[i]] + // inputs.push(input) + // vecEl.appendChild(input) + // } + + const widget = node.addDOMWidget(name, 'vector', inputEl, { + getValue() { + return JSON.stringify(widget._value) + }, + setValue(v) { + widget._value = v + }, + }) + const vec = new NumberInputWidget('vecEl', vector_size, true) + vec.setValue(...Object.values(value)) + vec.onChange = (value) => { + for (let i = 0; i < value.length; i++) { + const val = value[i] + widget._value[VECTOR_AXIS[i]] = Number.parseFloat(val) + } + + widget.callback?.(widget._value) + // widget._value[VECTOR_AXIS[index]] = Number.parseFloat(value) + } + + console.log('prev callback', widget.callback) + widget.callback = callback + widget._value = value + + // for (let i = 0; i < vector_size; i++) { + // const input = inputs[i] + // input.addEventListener('change', (event) => { + // widget._value[VECTOR_AXIS[i]] = Number.parseFloat(event.target.value) + // widget.callback?.(widget._value) + // node.graph._version++ + // node.setDirtyCanvas(true, true) + // }) + // } + + widget.inputEl = inputEl + widget.vecEl = vecEl + widget.vec = vec + + return { minWidth: 400, minHeight: 200 * vector_size, widget } +} export const MtbWidgets = { + //TODO: complete this properly + + /** + * Creates a vector widget. + * @param {string} key - The key for the widget. + * @param {number[]} [val] - The initial value for the widget. + * @param {number} size - The size of the vector. + * @returns {VectorWidget} The vector widget. + */ + VECTOR: (key, val, size) => { + shared.infoLogger('Adding VECTOR widget', { key, val, size }) + /** @type {VectorWidget} */ + const widget = { + name: key, + type: `vector${size}`, + y: 0, + options: { default: Array.from({ length: size }, () => 0.0) }, + _value: val || Array.from({ length: size }, () => 0.0), + draw: function (ctx, node, width, widgetY, height) { + ctx.textAlign = 'left' + ctx.strokeStyle = outline_color + ctx.fillStyle = background_color + ctx.beginPath() + if (show_text) + ctx.roundRect(margin, y, widget_width - margin * 2, H, [H * 0.5]) + else ctx.rect(margin, y, widget_width - margin * 2, H) + ctx.fill() + if (show_text) { + if (!w.disabled) ctx.stroke() + ctx.fillStyle = text_color + if (!w.disabled) { + ctx.beginPath() + ctx.moveTo(margin + 16, y + 5) + ctx.lineTo(margin + 6, y + H * 0.5) + ctx.lineTo(margin + 16, y + H - 5) + ctx.fill() + ctx.beginPath() + ctx.moveTo(widget_width - margin - 16, y + 5) + ctx.lineTo(widget_width - margin - 6, y + H * 0.5) + ctx.lineTo(widget_width - margin - 16, y + H - 5) + ctx.fill() + } + ctx.fillStyle = secondary_text_color + ctx.fillText(w.label || w.name, margin * 2 + 5, y + H * 0.7) + ctx.fillStyle = text_color + ctx.textAlign = 'right' + if (w.type === 'number') { + ctx.fillText( + Number(w.value).toFixed( + w.options.precision !== undefined ? w.options.precision : 3, + ), + widget_width - margin * 2 - 20, + y + H * 0.7, + ) + } else { + let v = w.value + if (w.options.values) { + let values = w.options.values + if (values.constructor === Function) values = values() + if (values && values.constructor !== Array) v = values[w.value] + } + ctx.fillText(v, widget_width - margin * 2 - 20, y + H * 0.7) + } + } + }, + get value() { + return this._value + }, + set value(val) { + this._value = val + this.callback?.(this._value) + }, + } + + return widget + }, BBOX: (key, val) => { /** @type {import("./types/litegraph").IWidget} */ const widget = { @@ -460,14 +708,8 @@ const mtb_widgets = { }, }) }, - registerCustomNodes() { - LiteGraph.registerNodeType('Constant (mtb)', Constant) - Constant.category = 'mtb/utils' - Constant.title = 'Constant (mtb)' - }, - - getCustomWidgets: function () { + getCustomWidgets: () => { return { BOOL: (node, inputName, inputData, app) => { console.debug('Registering bool') @@ -928,11 +1170,9 @@ const mtb_widgets = { const r = onConnectionsChange ? onConnectionsChange.apply(this, arguments) : undefined - shared.dynamic_connection(this, index, connected, 'var_', '*', [ - 'x', - 'y', - 'z', - ]) + shared.dynamic_connection(this, index, connected, 'var_', '*', { + nameArray: ['x', 'y', 'z'], + }) //- infer type if (link_info) {