From b9056441fe0010cce6fae13391d6970b91c08334 Mon Sep 17 00:00:00 2001 From: unknown Date: Wed, 23 Apr 2025 19:36:37 +0800 Subject: [PATCH] process bar is ok --- RH_ExecuteNode.py | 249 ++++++++++++++++++++++++++++------------------ 1 file changed, 153 insertions(+), 96 deletions(-) diff --git a/RH_ExecuteNode.py b/RH_ExecuteNode.py index cfbba84..465b2ba 100644 --- a/RH_ExecuteNode.py +++ b/RH_ExecuteNode.py @@ -31,15 +31,49 @@ class ExecuteNode: self.pbar = None self.node_lock = threading.Lock() self.total_nodes = None - # 添加当前步数追踪 - self.current_steps = 0 + self.current_steps = 0 # Track current steps for logging + + def update_progress(self): + """Increments the progress bar by one step and logs, stopping at total_nodes.""" + # --- Guard Condition --- + # Use lock to ensure thread safety when checking/updating steps and flag + with self.node_lock: + if self.task_completed or (self.pbar and self.current_steps >= self.total_nodes): + # Print only if trying to update *after* completion for debugging + if self.task_completed: + print(f"Skipping progress update because task is already completed.") + return + + if self.pbar: + self.current_steps += 1 + # Increment the ComfyUI progress bar by 1 + self.pbar.update(1) + # Log the current state + display_steps = min(self.current_steps, self.total_nodes) # Ensure log doesn't exceed total + print(f"Progress Update: Step {display_steps}/{self.total_nodes} ({(display_steps/self.total_nodes)*100:.1f}%)") + + + def complete_progress(self): + """Sets the progress bar to 100% and marks task as completed.""" + # --- Use lock for thread safety --- + with self.node_lock: + # Check if already completed to prevent redundant calls/logs + if self.task_completed: + return + + print(f"Finalizing progress: Setting task_completed = True") + # --- Set completion flag FIRST --- + self.task_completed = True + + if self.pbar: + # Ensure internal step count reflects completion + self.current_steps = self.total_nodes + # Use update_absolute for the final state + self.pbar.update_absolute(1.0) + print(f"Progress Finalized: {self.total_nodes}/{self.total_nodes} (100.0%)") + else: + print("Progress bar not available during finalization.") - def update(self, step_num): - """直接更新进度条步数,与EasyControlGenerate保持一致""" - if self.pbar: - # 传递当前步数,让ProgressBar自己处理百分比计算 - self.pbar.update(step_num) - print(f"Progress: {step_num}/{self.total_nodes}") @classmethod def INPUT_TYPES(cls): @@ -68,6 +102,12 @@ class ExecuteNode: def on_ws_message(self, ws, message): """处理 WebSocket 消息,更新内部状态和进度条""" try: + # --- Check completion status AT THE START --- + # This check is implicitly thread-safe due to complete_progress lock + if self.task_completed: + # print("WS Message received after task completion, ignoring.") # Optional: reduce log spam + return + data = json.loads(message) message_type = data.get("type") @@ -75,51 +115,47 @@ class ExecuteNode: node_data = data.get("data", {}) node_id = node_data.get("node") if node_id is not None: - with self.node_lock: - if node_id not in self.executed_nodes: - self.executed_nodes.add(node_id) - # 调用update并传递当前节点数 - current_progress = len(self.executed_nodes) - # 关键:使用直接的步数更新,完全模仿EasyControlGenerate - self.update(current_progress) - print(f"WS: Node {node_id} executed. Progress: {current_progress}/{self.total_nodes}") + # Lock is handled within update_progress now + # Check if it's a new node before calling update + if node_id not in self.executed_nodes: + self.executed_nodes.add(node_id) # Add before update call + self.update_progress() # This method is now guarded internally + print(f"WS: Node {node_id} executed.") else: - if not self.task_completed: - print("WS: Received null node signal, assuming task nearing completion...") - # 关键修改:设置一个较高的进度,但不设为完成 - # 设为90%,保留10%给最终成功信号 - self.update(int(self.total_nodes * 0.9)) - # 启动一个计时器,如果5秒内没有收到执行成功信号,则强制标记为完成 - threading.Timer(5.0, self.check_and_complete_task).start() + # Null node signal check remains guarded by the top-level check + print("WS: Received null node signal, waiting for final success signal...") elif message_type == "execution_success": - if not self.task_completed: - print("WS: Task execution success signal received.") - # 确保进度条在任务完成时处于100% - # 重要:需要调用update传递总步数 - self.update(self.total_nodes) - time.sleep(2) # 给服务器一点时间完成最终处理 - self.task_completed = True + # The internal check in complete_progress handles redundancy + print("WS: Task execution success signal received.") + self.complete_progress() + # No need for sleep or setting task_completed here except Exception as e: print(f"Error processing WebSocket message: {e}") self.ws_error = e - self.task_completed = True + # Call complete_progress which handles the task_completed flag and lock + self.complete_progress() def on_ws_error(self, ws, error): """处理 WebSocket 错误""" print(f"WebSocket error: {error}") self.ws_error = error - self.task_completed = True # Assume task is over + # Mark task as complete via the centralized method + self.complete_progress() def on_ws_close(self, ws, close_status_code, close_msg): """处理 WebSocket 关闭""" print(f"WebSocket closed: {close_status_code} - {close_msg}") # If closed unexpectedly, mark as complete to end loop - if not self.task_completed: - print("Warning: WebSocket closed unexpectedly.") + # Use lock temporarily just to read task_completed safely + with self.node_lock: + should_complete = not self.task_completed + if should_complete: + print("Warning: WebSocket closed unexpectedly. Forcing task completion.") self.ws_error = self.ws_error or IOError(f"WebSocket closed unexpectedly ({close_status_code})") - self.task_completed = True + # Mark task as complete via the centralized method + self.complete_progress() def on_ws_open(self, ws): """处理 WebSocket 连接打开""" @@ -143,43 +179,38 @@ class ExecuteNode: print("WebSocket thread started.") def check_and_complete_task(self): - """如果任务仍未完成,强制标记为完成""" - if not self.task_completed: - print("Task completion timeout after null node signal - forcing completion.") - self.update(self.total_nodes) - self.task_completed = True + """If task times out after null node, force completion.""" + # complete_progress now checks the flag internally and uses lock + print("Task completion timeout after null node signal - attempting forced completion.") + self.complete_progress() # --- Main Process Method --- - # Removed query_interval from signature as it's no longer used for loop timing def process(self, apiConfig, nodeInfoList=None, run_timeout=600, concurrency_limit=1, estimated_total_nodes=ESTIMATED_TOTAL_NODES): - # 1. Reset State - with self.node_lock: + # Reset state + with self.node_lock: # Use lock for resetting shared state self.executed_nodes.clear() - self.task_completed = False - self.ws_error = None - self.prompt_tips = "{}" - - # 重置当前步数 - self.current_steps = 0 - - # 设置总节点数 + self.task_completed = False + self.ws_error = None + self.prompt_tips = "{}" + self.current_steps = 0 # Reset step counter + + # 设置总节点数并初始化进度条 self.total_nodes = max(1, estimated_total_nodes) print(f"Using total nodes for progress: {self.total_nodes}") - # 使用初始化和更新方式与EasyControlGenerate完全一致 + # Initialize ComfyUI progress bar self.pbar = comfy.utils.ProgressBar(self.total_nodes) - self.update(0) # 初始化进度为0 print("Progress bar initialized at 0") - # 3. Setup & Pre-checks (Concurrency) - print(f"Concurrency limit set to: {concurrency_limit}") + # --- Concurrency Check --- + api_key = None + base_url = None try: - # Use safer .get() access for apiConfig api_key = apiConfig.get("apiKey") base_url = apiConfig.get("base_url") if not api_key or not base_url: raise ValueError("apiKey and base_url missing from apiConfig") - + account_status = self.check_account_status(api_key, base_url) current_tasks = int(account_status["currentTaskCounts"]) print(f"There are {current_tasks} tasks running") @@ -191,7 +222,7 @@ class ExecuteNode: wait_interval = 2 # seconds while current_tasks >= concurrency_limit: if time.time() - start_wait_time > run_timeout: - if self.pbar: self.pbar.update_absolute(1.0) # Finish bar + if self.pbar: self.pbar.update_absolute(1.0) # Use absolute directly for setup failure raise Exception(f"Timeout waiting for concurrent tasks ({current_tasks}/{concurrency_limit}) to finish.") print(f"Waiting for concurrent tasks... ({current_tasks}/{concurrency_limit})") time.sleep(wait_interval) @@ -200,11 +231,11 @@ class ExecuteNode: print("Concurrency slot available.") except Exception as e: print(f"Error checking account status or waiting: {e}") - if self.pbar: self.pbar.update_absolute(1.0) # Ensure bar finishes + if self.pbar: self.pbar.update_absolute(1.0) # Use absolute directly for setup failure raise - # 4. Create Task & Connect WebSocket - task_id = None # Ensure task_id is defined for finally block + # --- Task Creation & WebSocket --- + task_id = None try: print(f"ExecuteNode NodeInfoList: {nodeInfoList}") # Pass base_url explicitly from the validated config @@ -223,65 +254,91 @@ class ExecuteNode: raise ValueError("Missing taskId or netWssUrl in task creation response.") print(f"Task created successfully, taskId: {task_id}") - self.connect_websocket(wss_url) + self.connect_websocket(wss_url) except Exception as e: print(f"Error creating task or connecting WS: {e}") - if self.pbar: self.pbar.update_absolute(1.0) + if self.pbar: self.pbar.update_absolute(1.0) # Use absolute directly for setup failure raise - # 5. Task Monitoring Loop (Focus on WS state and pbar update) + # --- Task Monitoring Loop --- task_start_time = time.time() - # Shorter sleep interval for more responsive UI updates - loop_sleep_interval = 0.1 + loop_sleep_interval = 0.1 print("Starting task monitoring loop...") - # 添加一个全局超时以防止永久卡死,无论发生什么情况 timeout_timer = None try: - # 设置一个全局超时计时器 + # Setup global timeout timer def force_timeout(): - if not self.task_completed: + # Use lock to safely check task_completed + with self.node_lock: + is_completed = self.task_completed + if not is_completed: print("Global timeout reached - forcing task completion.") self.ws_error = Exception("Global timeout reached") - self.task_completed = True - + # Let the main loop call complete_progress via the finally block or error handling + # Just set the flags here to break loop + self.task_completed = True # Set flag directly here to break loop + timeout_timer = threading.Timer(run_timeout, force_timeout) timeout_timer.daemon = True timeout_timer.start() - - # 使用带有超时的等待循环 - wait_start_time = time.time() - while not self.task_completed and not self.ws_error: - # 检查总超时 - if time.time() - wait_start_time > run_timeout: - print("Task monitoring timeout reached.") - self.ws_error = Exception(f"Timeout: Task {task_id} did not complete within {run_timeout} seconds.") - self.task_completed = True - break - - # 短暂睡眠以避免高CPU使用率 - time.sleep(0.1) - - # 处理退出条件 - if self.ws_error: - print(f"Task ended with error: {self.ws_error}") - raise self.ws_error - - print("Task monitoring completed successfully.") - + + # Main wait loop + while True: + # Check completion flags (read safely with lock) + with self.node_lock: + is_completed = self.task_completed + current_error = self.ws_error + if is_completed or current_error: + break # Exit loop if completed or error occurred + + # Check for timeout explicitly in loop as backup/alternative to timer + if time.time() - task_start_time > run_timeout: + print("Task monitoring loop timeout check triggered.") + # Set flags to exit loop; rely on finally block for completion + with self.node_lock: + if not self.task_completed: # Avoid overwriting WS error + self.ws_error = self.ws_error or Exception(f"Timeout: Task {task_id} did not complete within {run_timeout} seconds.") + self.task_completed = True # Ensure loop exit + break # Exit loop + + time.sleep(loop_sleep_interval) # Yield CPU + + # Handle exit conditions after loop + with self.node_lock: # Read error flag safely + final_error = self.ws_error + + if final_error: + print(f"Task ended with error: {final_error}") + # Only complete progress if not already completed by WS handler + with self.node_lock: + if not self.task_completed: + self.complete_progress() + raise final_error # Re-raise the error + else: # Task completed normally + print("Task monitoring completed successfully.") + # complete_progress should have been called by WS handler + finally: - # 清理超时计时器 + # Cleanup if timeout_timer: timeout_timer.cancel() - # 确保WebSocket已关闭 if self.ws: try: self.ws.close() except Exception as e: print(f"Error closing WebSocket: {e}") self.ws = None - - # 7. Process Output + + # Final safety net: Only complete progress if not already completed + with self.node_lock: + is_finally_completed = self.task_completed + if not is_finally_completed: + print("Warning: Monitoring loop ended unexpectedly. Finalizing progress via finally block.") + self.complete_progress() # Call the safe completion method + + + # --- Process Output --- print("Processing task output...") # Pass the validated api_key and base_url again return self.process_task_output(task_id, api_key, base_url)