CoolFace
Apppublic

Hamza4100/ai-workflow-agent

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
comfyui_builder.py547 linesDownload Raw Back to tools
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