diff --git a/__init__.py b/__init__.py index 7f87b7f..8777901 100644 --- a/__init__.py +++ b/__init__.py @@ -4,4 +4,6 @@ Diffusion Self-Distillation ComfyUI nodes from .dsd_nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file +WEB_DIRECTORY = "./web" + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] \ No newline at end of file diff --git a/dsd_nodes.py b/dsd_nodes.py index eccbcd9..cbc3655 100644 --- a/dsd_nodes.py +++ b/dsd_nodes.py @@ -145,14 +145,25 @@ class DSDGeminiPromptEnhancer: RETURN_NAMES = ("enhanced_prompt",) FUNCTION = "enhance_prompt" CATEGORY = "DSD" + OUTPUT_NODE = True # This ensures that UI data is sent to the node + + def __init__(self): + self.enhanced_prompt = None + + def get_state(self): + return { + "enhanced_prompt": self.enhanced_prompt + } def enhance_prompt(self, image, prompt, api_key): if not IMPORTS_AVAILABLE: print("Warning: DSD modules not available. Using original prompt.") + self.enhanced_prompt = None return (prompt,) if not GEMINI_AVAILABLE: print("Warning: Google Gemini API not available. Returning original prompt.") + self.enhanced_prompt = None return (prompt,) if not api_key: @@ -160,6 +171,7 @@ class DSDGeminiPromptEnhancer: api_key = os.getenv("GEMINI_API_KEY") if not api_key: print("Warning: No API key provided for Gemini. Returning original prompt.") + self.enhanced_prompt = None return (prompt,) # Convert from ComfyUI image to PIL @@ -169,14 +181,20 @@ class DSDGeminiPromptEnhancer: try: # Call the imported enhance_prompt function - enhanced_prompt = enhance_prompt(pil_image, prompt,api_key) + enhanced_prompt = enhance_prompt(pil_image, prompt, api_key) print("Original prompt:", prompt) print("Enhanced prompt:", enhanced_prompt) - return (enhanced_prompt,) + # Store the enhanced prompt for UI display + self.enhanced_prompt = enhanced_prompt + + # Return the enhanced prompt and explicitly include it in the UI data + # Make sure enhanced_prompt is a proper string, not an array/list of characters + return {"ui": {"enhanced_prompt": str(enhanced_prompt)}, "result": (enhanced_prompt,)} except Exception as e: print(f"Error enhancing prompt: {e}") + self.enhanced_prompt = None return (prompt,) diff --git a/web/js/showEnhancedPrompt.js b/web/js/showEnhancedPrompt.js new file mode 100644 index 0000000..f9be3e6 --- /dev/null +++ b/web/js/showEnhancedPrompt.js @@ -0,0 +1,77 @@ +import { app } from "../../../scripts/app.js"; +import { ComfyWidgets } from "../../../scripts/widgets.js"; + +// Displays enhanced prompt on DSDGeminiPromptEnhancer node +app.registerExtension({ + name: "comfyui.dsd.showEnhancedPrompt", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData.name === "DSDGeminiPromptEnhancer") { + function normalizePrompt(prompt) { + if (!prompt) return ""; + if (Array.isArray(prompt)) { + return prompt.join(""); + } + return prompt; + } + + function populateEnhancedPrompt(enhancedPrompt) { + enhancedPrompt = normalizePrompt(enhancedPrompt); + if (!enhancedPrompt) { + return; + } + + try { + if (this.widgets) { + for (let i = 0; i < this.widgets.length; i++) { + if (this.widgets[i].name === "enhanced_prompt_display") { + this.widgets[i].onRemove?.(); + this.widgets.splice(i, 1); + i--; + } + } + } + + const textWidgetResult = ComfyWidgets["STRING"](this, "enhanced_prompt_text", ["STRING", { multiline: true }], app); + if (!textWidgetResult || !textWidgetResult.widget || !textWidgetResult.widget.inputEl) { + return; + } + const w = textWidgetResult.widget; + w.inputEl.readOnly = true; + w.inputEl.style.opacity = 0.8; + w.inputEl.style.backgroundColor = "#1e2124"; + w.inputEl.style.color = "#9eec51"; + w.name = "enhanced_prompt_display"; + w.value = enhancedPrompt; + + requestAnimationFrame(() => { + const sz = this.computeSize(); + if (sz[0] < this.size[0]) sz[0] = this.size[0]; + if (sz[1] < this.size[1]) sz[1] = this.size[1]; + this.onResize?.(sz); + app.graph.setDirtyCanvas(true, false); + }); + } catch (error) { + // Silent error handling + } + } + + const onExecuted = nodeType.prototype.onExecuted; + nodeType.prototype.onExecuted = function (message) { + try { + onExecuted?.apply(this, arguments); + if (message.enhanced_prompt) { + populateEnhancedPrompt.call(this, message.enhanced_prompt); + } else if (message.ui && message.ui.enhanced_prompt) { + populateEnhancedPrompt.call(this, message.ui.enhanced_prompt); + } else if (message.text) { + populateEnhancedPrompt.call(this, message.text); + } + } catch (error) { + // Silent error handling + } + }; + } + }, +}); + +window.DSD_ENHANCED_PROMPT_LOADED = true; \ No newline at end of file