Hamza4100/ai-workflow-agent
0
1# ComfyUI Workflow Builder Tool
2"""
3Generate and execute ComfyUI workflow JSON templates.
4Supports common generative AI patterns.
5LLM-enhanced generation when context.use_llm=True.
6"""
7
8import httpx
9import json
10import logging
11import uuid
12from typing import Dict, Any, List, Optional
13
14from config import settings
15
16logger = logging.getLogger(__name__)
17
18
19class ComfyUIWorkflowBuilder:
20 """
21 ComfyUI workflow generator and executor.
22 Creates JSON workflow graphs and executes via ComfyUI API.
23 """
24
25 def __init__(self):
26 self.comfyui_host = settings.COMFYUI_HOST
27 self.client = httpx.AsyncClient(timeout=300.0) # Long timeout for image generation
28
29 async def check_health(self) -> str:
30 """Check if ComfyUI is running and responsive."""
31 try:
32 response = await self.client.get(f"{self.comfyui_host}/system_stats")
33 if response.status_code == 200:
34 return "healthy"
35 return "unhealthy"
36 except Exception as e:
37 logger.debug(f"ComfyUI health check failed: {e}")
38 return "unreachable"
39
40 async def generate_workflow(self, query: str, context: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
41 """
42 Generate ComfyUI workflow JSON based on user query.
43
44 Args:
45 query: User's natural language request
46 context: Optional context with use_llm flag for LLM-enhanced generation
47
48 Returns:
49 ComfyUI workflow JSON structure
50 """
51 # Check if LLM enhancement is requested
52 use_llm = False
53 if context and isinstance(context, dict):
54 use_llm = context.get('use_llm', False)
55
56 # If LLM mode is enabled, use AI to analyze and create more intelligent workflow
57 if use_llm:
58 logger.info("Using LLM-enhanced workflow generation")
59 workflow = await self._generate_llm_workflow(query)
60 if workflow:
61 return workflow
62 # Fall back to template if LLM fails
63 logger.warning("LLM generation failed, falling back to templates")
64
65 # Template-based generation (keyword mode)
66 workflow_type = self._detect_workflow_type(query)
67
68 # Extract parameters from query
69 params = self._extract_params(query)
70
71 # Generate appropriate template
72 templates = {
73 "txt2img": self._generate_txt2img_workflow,
74 "img2img": self._generate_img2img_workflow,
75 "upscale": self._generate_upscale_workflow,
76 "inpaint": self._generate_inpaint_workflow,
77 "controlnet": self._generate_controlnet_workflow,
78 "generic": self._generate_generic_workflow
79 }
80
81 generator = templates.get(workflow_type, self._generate_generic_workflow)
82 workflow = generator(params)
83
84 return workflow
85
86 async def _generate_llm_workflow(self, query: str) -> Optional[Dict[str, Any]]:
87 """
88 Use LLM to generate a more intelligent workflow based on query analysis.
89
90 Args:
91 query: User's natural language request
92
93 Returns:
94 Enhanced ComfyUI workflow or None if LLM fails
95 """
96 try:
97 # Import here to avoid circular dependency
98 from decision_agent import DecisionAgent
99
100 agent = DecisionAgent()
101 analysis = await agent.analyze(query, context={'use_llm': True})
102
103 # Use analysis explanation to create more detailed workflow
104 workflow_type = self._detect_workflow_type(query)
105 explanation = analysis.get('explanation', '')
106 params = self._extract_params(query)
107
108 # Generate base template
109 templates = {
110 "txt2img": self._generate_txt2img_workflow,
111 "img2img": self._generate_img2img_workflow,
112 "upscale": self._generate_upscale_workflow,
113 "inpaint": self._generate_inpaint_workflow,
114 "controlnet": self._generate_controlnet_workflow,
115 "generic": self._generate_generic_workflow
116 }
117
118 generator = templates.get(workflow_type, self._generate_generic_workflow)
119 workflow = generator(params)
120
121 # Enhance with LLM analysis
122 workflow['meta']['llm_analysis'] = {
123 'explanation': explanation,
124 'confidence': analysis.get('confidence', 0.0),
125 'suggested_tools': analysis.get('suggested_tools', []),
126 'next_steps': analysis.get('next_steps', [])
127 }
128 workflow['meta']['generated_with_llm'] = True
129
130 return workflow
131
132 except Exception as e:
133 logger.error(f"LLM workflow generation failed: {e}")
134 return None
135
136 def _detect_workflow_type(self, query: str) -> str:
137 """Detect the type of ComfyUI workflow needed."""
138 query_lower = query.lower()
139
140 if any(w in query_lower for w in ["upscale", "enhance", "higher resolution", "4x", "2x"]):
141 return "upscale"
142 elif any(w in query_lower for w in ["inpaint", "edit", "remove", "fill", "mask"]):
143 return "inpaint"
144 elif any(w in query_lower for w in ["controlnet", "pose", "depth", "canny", "edge"]):
145 return "controlnet"
146 elif any(w in query_lower for w in ["img2img", "transform", "style transfer", "from image"]):
147 return "img2img"
148 else:
149 return "txt2img"
150
151 def _extract_params(self, query: str) -> Dict[str, Any]:
152 """Extract generation parameters from query."""
153 # Default parameters
154 params = {
155 "prompt": query,
156 "negative_prompt": "bad quality, blurry, deformed",
157 "width": 512,
158 "height": 512,
159 "steps": 20,
160 "cfg": 7.0,
161 "seed": -1, # Random
162 "checkpoint": "v1-5-pruned-emaonly.safetensors"
163 }
164
165 query_lower = query.lower()
166
167 # Detect resolution
168 if "portrait" in query_lower or "vertical" in query_lower:
169 params["width"] = 512
170 params["height"] = 768
171 elif "landscape" in query_lower or "horizontal" in query_lower:
172 params["width"] = 768
173 params["height"] = 512
174 elif "square" in query_lower:
175 params["width"] = 512
176 params["height"] = 512
177 elif "hd" in query_lower or "1024" in query_lower:
178 params["width"] = 1024
179 params["height"] = 1024
180
181 # Detect model
182 if "sdxl" in query_lower:
183 params["checkpoint"] = "sd_xl_base_1.0.safetensors"
184 params["width"] = 1024
185 params["height"] = 1024
186 elif "flux" in query_lower:
187 params["checkpoint"] = "flux1-dev.safetensors"
188
189 # Detect quality settings
190 if "high quality" in query_lower or "detailed" in query_lower:
191 params["steps"] = 30
192 params["cfg"] = 8.0
193 elif "fast" in query_lower or "quick" in query_lower:
194 params["steps"] = 15
195 params["cfg"] = 6.0
196
197 return params
198
199 def _generate_txt2img_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
200 """Generate text-to-image workflow."""
201 return {
202 "prompt": {
203 "3": {
204 "inputs": {
205 "seed": params.get("seed", -1),
206 "steps": params.get("steps", 20),
207 "cfg": params.get("cfg", 7.0),
208 "sampler_name": "euler",
209 "scheduler": "normal",
210 "denoise": 1.0,
211 "model": ["4", 0],
212 "positive": ["6", 0],
213 "negative": ["7", 0],
214 "latent_image": ["5", 0]
215 },
216 "class_type": "KSampler",
217 "_meta": {"title": "KSampler"}
218 },
219 "4": {
220 "inputs": {
221 "ckpt_name": params.get("checkpoint", "v1-5-pruned-emaonly.safetensors")
222 },
223 "class_type": "CheckpointLoaderSimple",
224 "_meta": {"title": "Load Checkpoint"}
225 },
226 "5": {
227 "inputs": {
228 "width": params.get("width", 512),
229 "height": params.get("height", 512),
230 "batch_size": 1
231 },
232 "class_type": "EmptyLatentImage",
233 "_meta": {"title": "Empty Latent Image"}
234 },
235 "6": {
236 "inputs": {
237 "text": params.get("prompt", "beautiful landscape"),
238 "clip": ["4", 1]
239 },
240 "class_type": "CLIPTextEncode",
241 "_meta": {"title": "CLIP Text Encode (Prompt)"}
242 },
243 "7": {
244 "inputs": {
245 "text": params.get("negative_prompt", "bad quality, blurry"),
246 "clip": ["4", 1]
247 },
248 "class_type": "CLIPTextEncode",
249 "_meta": {"title": "CLIP Text Encode (Negative)"}
250 },
251 "8": {
252 "inputs": {
253 "samples": ["3", 0],
254 "vae": ["4", 2]
255 },
256 "class_type": "VAEDecode",
257 "_meta": {"title": "VAE Decode"}
258 },
259 "9": {
260 "inputs": {
261 "filename_prefix": "ComfyUI",
262 "images": ["8", 0]
263 },
264 "class_type": "SaveImage",
265 "_meta": {"title": "Save Image"}
266 }
267 },
268 "meta": {
269 "generated_by": "AI Workflow Agent",
270 "type": "txt2img",
271 "params": params
272 }
273 }
274
275 def _generate_img2img_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
276 """Generate image-to-image workflow."""
277 workflow = self._generate_txt2img_workflow(params)
278
279 # Modify for img2img
280 workflow["prompt"]["5"] = {
281 "inputs": {
282 "image": "INPUT_IMAGE_PATH",
283 "upload": "image"
284 },
285 "class_type": "LoadImage",
286 "_meta": {"title": "Load Image"}
287 }
288
289 # Add VAE encode for input
290 workflow["prompt"]["10"] = {
291 "inputs": {
292 "pixels": ["5", 0],
293 "vae": ["4", 2]
294 },
295 "class_type": "VAEEncode",
296 "_meta": {"title": "VAE Encode"}
297 }
298
299 # Update sampler to use encoded image
300 workflow["prompt"]["3"]["inputs"]["latent_image"] = ["10", 0]
301 workflow["prompt"]["3"]["inputs"]["denoise"] = 0.75
302
303 workflow["meta"]["type"] = "img2img"
304
305 return workflow
306
307 def _generate_upscale_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
308 """Generate upscale workflow."""
309 return {
310 "prompt": {
311 "1": {
312 "inputs": {
313 "image": "INPUT_IMAGE_PATH",
314 "upload": "image"
315 },
316 "class_type": "LoadImage",
317 "_meta": {"title": "Load Image"}
318 },
319 "2": {
320 "inputs": {
321 "model_name": "RealESRGAN_x4plus.pth"
322 },
323 "class_type": "UpscaleModelLoader",
324 "_meta": {"title": "Load Upscale Model"}
325 },
326 "3": {
327 "inputs": {
328 "upscale_model": ["2", 0],
329 "image": ["1", 0]
330 },
331 "class_type": "ImageUpscaleWithModel",
332 "_meta": {"title": "Upscale Image"}
333 },
334 "4": {
335 "inputs": {
336 "filename_prefix": "Upscaled",
337 "images": ["3", 0]
338 },
339 "class_type": "SaveImage",
340 "_meta": {"title": "Save Image"}
341 }
342 },
343 "meta": {
344 "generated_by": "AI Workflow Agent",
345 "type": "upscale",
346 "params": params
347 }
348 }
349
350 def _generate_inpaint_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
351 """Generate inpainting workflow."""
352 workflow = self._generate_txt2img_workflow(params)
353
354 # Add mask loading
355 workflow["prompt"]["10"] = {
356 "inputs": {
357 "image": "INPUT_IMAGE_PATH",
358 "upload": "image"
359 },
360 "class_type": "LoadImage",
361 "_meta": {"title": "Load Image"}
362 }
363
364 workflow["prompt"]["11"] = {
365 "inputs": {
366 "image": "MASK_IMAGE_PATH",
367 "upload": "image"
368 },
369 "class_type": "LoadImage",
370 "_meta": {"title": "Load Mask"}
371 }
372
373 # Replace empty latent with masked image
374 workflow["prompt"]["5"] = {
375 "inputs": {
376 "grow_mask_by": 6,
377 "pixels": ["10", 0],
378 "vae": ["4", 2],
379 "mask": ["11", 0]
380 },
381 "class_type": "VAEEncodeForInpaint",
382 "_meta": {"title": "VAE Encode (Inpaint)"}
383 }
384
385 workflow["meta"]["type"] = "inpaint"
386
387 return workflow
388
389 def _generate_controlnet_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
390 """Generate ControlNet workflow."""
391 workflow = self._generate_txt2img_workflow(params)
392
393 # Add ControlNet
394 workflow["prompt"]["10"] = {
395 "inputs": {
396 "control_net_name": "control_v11p_sd15_canny.pth"
397 },
398 "class_type": "ControlNetLoader",
399 "_meta": {"title": "Load ControlNet"}
400 }
401
402 workflow["prompt"]["11"] = {
403 "inputs": {
404 "image": "CONTROL_IMAGE_PATH",
405 "upload": "image"
406 },
407 "class_type": "LoadImage",
408 "_meta": {"title": "Load Control Image"}
409 }
410
411 workflow["prompt"]["12"] = {
412 "inputs": {
413 "strength": 1.0,
414 "conditioning": ["6", 0],
415 "control_net": ["10", 0],
416 "image": ["11", 0]
417 },
418 "class_type": "ControlNetApply",
419 "_meta": {"title": "Apply ControlNet"}
420 }
421
422 # Update sampler to use ControlNet conditioning
423 workflow["prompt"]["3"]["inputs"]["positive"] = ["12", 0]
424
425 workflow["meta"]["type"] = "controlnet"
426
427 return workflow
428
429 def _generate_generic_workflow(self, params: Dict[str, Any]) -> Dict[str, Any]:
430 """Generate generic workflow (defaults to txt2img)."""
431 return self._generate_txt2img_workflow(params)
432
433 async def execute_workflow(self, workflow: Dict[str, Any]) -> Dict[str, Any]:
434 """
435 Execute workflow in ComfyUI.
436
437 Args:
438 workflow: ComfyUI workflow JSON
439
440 Returns:
441 Execution result with output paths
442 """
443 try:
444 # Get the prompt part of workflow
445 prompt = workflow.get("prompt", workflow)
446
447 # Generate client ID
448 client_id = str(uuid.uuid4())
449
450 # Queue the prompt
451 response = await self.client.post(
452 f"{self.comfyui_host}/prompt",
453 json={
454 "prompt": prompt,
455 "client_id": client_id
456 }
457 )
458
459 if response.status_code == 200:
460 result = response.json()
461 prompt_id = result.get("prompt_id")
462
463 logger.info(f"ComfyUI prompt queued: {prompt_id}")
464
465 # Wait for completion (poll history)
466 output = await self._wait_for_completion(prompt_id)
467
468 return {
469 "success": True,
470 "prompt_id": prompt_id,
471 "output": output
472 }
473 else:
474 logger.error(f"ComfyUI queue failed: {response.status_code}")
475 return {
476 "success": False,
477 "error": f"Queue failed: {response.status_code}"
478 }
479
480 except Exception as e:
481 logger.error(f"ComfyUI execute error: {e}")
482 return {
483 "success": False,
484 "error": str(e)
485 }
486
487 async def _wait_for_completion(
488 self,
489 prompt_id: str,
490 timeout: int = 300,
491 poll_interval: int = 2
492 ) -> Dict[str, Any]:
493 """Wait for ComfyUI prompt to complete."""
494 import asyncio
495
496 elapsed = 0
497 while elapsed < timeout:
498 try:
499 response = await self.client.get(
500 f"{self.comfyui_host}/history/{prompt_id}"
501 )
502
503 if response.status_code == 200:
504 history = response.json()
505 if prompt_id in history:
506 return history[prompt_id]
507
508 await asyncio.sleep(poll_interval)
509 elapsed += poll_interval
510
511 except Exception as e:
512 logger.warning(f"Poll error: {e}")
513 await asyncio.sleep(poll_interval)
514 elapsed += poll_interval
515
516 return {"status": "timeout", "elapsed": elapsed}
517
518 async def get_models(self) -> List[str]:
519 """Get available models in ComfyUI."""
520 try:
521 response = await self.client.get(
522 f"{self.comfyui_host}/object_info/CheckpointLoaderSimple"
523 )
524
525 if response.status_code == 200:
526 data = response.json()
527 models = data.get("CheckpointLoaderSimple", {}).get(
528 "input", {}
529 ).get("required", {}).get("ckpt_name", [[]])[0]
530 return models
531 return []
532
533 except Exception as e:
534 logger.error(f"Get models error: {e}")
535 return []
536
537 async def get_queue_status(self) -> Dict[str, Any]:
538 """Get current ComfyUI queue status."""
539 try:
540 response = await self.client.get(f"{self.comfyui_host}/queue")
541 if response.status_code == 200:
542 return response.json()
543 return {}
544 except Exception as e:
545 logger.error(f"Queue status error: {e}")
546 return {}
547 