From 6c282d3f1b36c6cdb5889152afbf3604359ee611 Mon Sep 17 00:00:00 2001 From: AEmotionStudio <163354043+AEmotionStudio@users.noreply.github.com> Date: Wed, 14 Jan 2026 16:33:32 -0800 Subject: [PATCH] refactor: Rename utils to discordsend_utils --- .gitattributes | 5 - discord_image_node.py | 2 +- discord_video_node.py | 4 +- discordsend_utils/__init__.py | 19 ++ discordsend_utils/discord_api.py | 408 ++++++++++++++++++++++++ discordsend_utils/github_integration.py | 137 ++++++++ discordsend_utils/logging_config.py | 43 +++ discordsend_utils/prompt_extractor.py | 240 ++++++++++++++ discordsend_utils/sanitizer.py | 269 ++++++++++++++++ discordsend_utils/workflow_builder.py | 146 +++++++++ scripts/auto_review.sh | 42 +++ tests/test_utils.py | 4 +- 12 files changed, 1309 insertions(+), 10 deletions(-) delete mode 100644 .gitattributes create mode 100644 discordsend_utils/__init__.py create mode 100644 discordsend_utils/discord_api.py create mode 100644 discordsend_utils/github_integration.py create mode 100644 discordsend_utils/logging_config.py create mode 100644 discordsend_utils/prompt_extractor.py create mode 100644 discordsend_utils/sanitizer.py create mode 100644 discordsend_utils/workflow_builder.py create mode 100755 scripts/auto_review.sh diff --git a/.gitattributes b/.gitattributes deleted file mode 100644 index a6b56c9..0000000 --- a/.gitattributes +++ /dev/null @@ -1,5 +0,0 @@ -images/*.png filter=lfs diff=lfs merge=lfs -text -images/*.jpg filter=lfs diff=lfs merge=lfs -text -images/*.jpeg filter=lfs diff=lfs merge=lfs -text -images/*.webp filter=lfs diff=lfs merge=lfs -text -images/*.gif filter=lfs diff=lfs merge=lfs -text \ No newline at end of file diff --git a/discord_image_node.py b/discord_image_node.py index 2a00106..8961eb5 100644 --- a/discord_image_node.py +++ b/discord_image_node.py @@ -17,7 +17,7 @@ from uuid import uuid4 from typing import Any, Union, List, Optional # Import shared utilities -from utils import ( +from discordsend_utils import ( sanitize_json_for_export, update_github_cdn_urls, extract_prompts_from_workflow, diff --git a/discord_video_node.py b/discord_video_node.py index 909d5a4..2c3f517 100644 --- a/discord_video_node.py +++ b/discord_video_node.py @@ -23,7 +23,7 @@ import functools import server # Import shared utilities -from utils import sanitize_json_for_export, update_github_cdn_urls, send_to_discord_with_retry +from discordsend_utils import sanitize_json_for_export, update_github_cdn_urls, send_to_discord_with_retry # Define cached decorator for local use def cached(max_size=None): """ @@ -48,7 +48,7 @@ def cached(max_size=None): # Try to import dependencies from nodes.py try: - from utils import ffmpeg_path, get_audio, hash_path, validate_path, requeue_workflow, \ + from discordsend_utils import ffmpeg_path, get_audio, hash_path, validate_path, requeue_workflow, \ gifski_path, calculate_file_hash, strip_path, try_download_video, is_url, \ imageOrLatent, BIGMAX, merge_filter_args, ENCODE_ARGS, floatOrInt from comfy.utils import ProgressBar diff --git a/discordsend_utils/__init__.py b/discordsend_utils/__init__.py new file mode 100644 index 0000000..8f4b378 --- /dev/null +++ b/discordsend_utils/__init__.py @@ -0,0 +1,19 @@ +""" +ComfyUI-DiscordSend Utility Package + +Shared utilities for Discord integration, sanitization, and GitHub CDN operations. +""" + +from .sanitizer import sanitize_json_for_export +from .github_integration import update_github_cdn_urls +from .prompt_extractor import extract_prompts_from_workflow +from .discord_api import DiscordWebhookClient, validate_webhook_url, send_to_discord_with_retry + +__all__ = [ + 'sanitize_json_for_export', + 'update_github_cdn_urls', + 'extract_prompts_from_workflow', + 'DiscordWebhookClient', + 'validate_webhook_url', + 'send_to_discord_with_retry', +] diff --git a/discordsend_utils/discord_api.py b/discordsend_utils/discord_api.py new file mode 100644 index 0000000..556d0f3 --- /dev/null +++ b/discordsend_utils/discord_api.py @@ -0,0 +1,408 @@ +""" +Discord API Utilities for ComfyUI-DiscordSend + +Provides a client for interacting with Discord webhooks and validation utilities. +""" + +import os +import re +import time +import logging +from io import BytesIO +from typing import Any, Dict, List, Optional, Tuple + +import requests + +# Get logger for this module +logger = logging.getLogger("comfyui_discordsend") + + +# Discord webhook URL patterns +WEBHOOK_URL_PATTERNS = [ + r"https?://(?:www\.)?discord(?:app)?\.com/api/webhooks/\d+/[\w-]+", + r"https?://(?:www\.)?discordapp\.com/api/webhooks/\d+/[\w-]+", +] + + +def validate_webhook_url(url: str) -> Tuple[bool, str]: + """ + Validate a Discord webhook URL. + + Args: + url: The webhook URL to validate + + Returns: + Tuple of (is_valid, message) + """ + if not url: + return False, "Webhook URL is empty" + + if not url.startswith("http"): + return False, "Webhook URL must start with http:// or https://" + + # Check against known patterns + for pattern in WEBHOOK_URL_PATTERNS: + if re.match(pattern, url, re.IGNORECASE): + return True, "Valid Discord webhook URL" + + # More lenient check + if "discord" in url.lower() and "webhook" in url.lower(): + return True, "Appears to be a Discord webhook URL" + + return False, "URL does not appear to be a valid Discord webhook URL" + + +def sanitize_webhook_for_logging(url: str) -> str: + """ + Sanitize a webhook URL for safe logging (hide the token portion). + + Args: + url: The webhook URL + + Returns: + Sanitized URL safe for logging + """ + if not url: + return "" + + # Pattern: https://discord.com/api/webhooks/{id}/{token} + match = re.match(r"(https?://[^/]+/api/webhooks/\d+/)(.+)", url) + if match: + return f"{match.group(1)}[REDACTED]" + + return "[REDACTED_WEBHOOK_URL]" + + +class DiscordWebhookClient: + """ + Client for sending messages and files to Discord via webhooks. + + Features: + - Automatic retry with exponential backoff + - Rate limit handling + - File size validation + - Error handling with sanitized logging + """ + + # Discord limits + MAX_FILE_SIZE = 25 * 1024 * 1024 # 25MB for regular users + MAX_FILES_PER_MESSAGE = 10 + MAX_MESSAGE_LENGTH = 2000 + + def __init__(self, webhook_url: str, max_retries: int = 3): + """ + Initialize the Discord webhook client. + + Args: + webhook_url: The Discord webhook URL + max_retries: Maximum number of retry attempts for failed requests + """ + self.webhook_url = webhook_url + self.max_retries = max_retries + self._validated = False + + def validate(self) -> Tuple[bool, str]: + """Validate the webhook URL.""" + is_valid, message = validate_webhook_url(self.webhook_url) + self._validated = is_valid + return is_valid, message + + def send_message( + self, + content: str = "", + files: Optional[List[Tuple[str, bytes, str]]] = None, + embeds: Optional[List[Dict]] = None + ) -> Tuple[bool, Dict[str, Any]]: + """ + Send a message to Discord. + + Args: + content: Text message content + files: List of (filename, file_bytes, content_type) tuples + embeds: List of embed dictionaries + + Returns: + Tuple of (success, response_data) + """ + if not self.webhook_url: + return False, {"error": "No webhook URL configured"} + + # Validate URL if not already done + if not self._validated: + is_valid, message = self.validate() + if not is_valid: + return False, {"error": message} + + # Truncate message if too long + if content and len(content) > self.MAX_MESSAGE_LENGTH: + content = content[:self.MAX_MESSAGE_LENGTH - 3] + "..." + + # Prepare request + data = {} + if content: + data["content"] = content + if embeds: + data["embeds"] = embeds + + # Prepare files + request_files = None + if files: + request_files = [] + for i, (filename, file_bytes, content_type) in enumerate(files): + # Validate file size + if len(file_bytes) > self.MAX_FILE_SIZE: + continue # Skip oversized files + + request_files.append( + (f"file{i}", (filename, BytesIO(file_bytes), content_type)) + ) + + if len(request_files) > self.MAX_FILES_PER_MESSAGE: + request_files = request_files[:self.MAX_FILES_PER_MESSAGE] + + # Send with retry logic + return self._send_with_retry(data, request_files) + + def send_file( + self, + file_path: str, + message: str = "", + additional_files: Optional[List[str]] = None + ) -> Tuple[bool, Dict[str, Any]]: + """ + Send a file from disk to Discord. + + Args: + file_path: Path to the file to send + message: Optional message to accompany the file + additional_files: Optional list of additional file paths + + Returns: + Tuple of (success, response_data) + """ + files = [] + + # Add main file + if os.path.exists(file_path): + file_size = os.path.getsize(file_path) + if file_size <= self.MAX_FILE_SIZE: + with open(file_path, "rb") as f: + filename = os.path.basename(file_path) + content_type = self._get_content_type(filename) + files.append((filename, f.read(), content_type)) + + # Add additional files + if additional_files: + for path in additional_files: + if os.path.exists(path): + file_size = os.path.getsize(path) + if file_size <= self.MAX_FILE_SIZE: + with open(path, "rb") as f: + filename = os.path.basename(path) + content_type = self._get_content_type(filename) + files.append((filename, f.read(), content_type)) + + if not files: + return False, {"error": "No valid files to send"} + + return self.send_message(content=message, files=files) + + def _send_with_retry( + self, + data: Dict, + files: Optional[List] = None + ) -> Tuple[bool, Dict[str, Any]]: + """Send request with retry logic.""" + last_error = None + + for attempt in range(self.max_retries): + try: + if files: + response = requests.post( + self.webhook_url, + data={"payload_json": str(data)} if data else None, + files=files, + timeout=60 + ) + else: + response = requests.post( + self.webhook_url, + json=data, + timeout=30 + ) + + # Handle rate limiting + if response.status_code == 429: + retry_after = response.json().get("retry_after", 1) + time.sleep(retry_after) + continue + + # Success + if response.status_code in [200, 204]: + try: + return True, response.json() if response.content else {} + except: + return True, {} + + # Client errors (don't retry) + if 400 <= response.status_code < 500: + return False, { + "error": f"Discord API error: {response.status_code}", + "details": response.text[:500] + } + + # Server errors (retry) + last_error = f"Discord API returned {response.status_code}" + + except requests.exceptions.Timeout: + last_error = "Request timed out" + except requests.exceptions.RequestException as e: + last_error = str(e) + + # Exponential backoff + if attempt < self.max_retries - 1: + time.sleep(2 ** attempt) + + return False, {"error": last_error or "Unknown error"} + + def _get_content_type(self, filename: str) -> str: + """Get content type based on file extension.""" + ext = os.path.splitext(filename)[1].lower() + content_types = { + ".png": "image/png", + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".gif": "image/gif", + ".webp": "image/webp", + ".mp4": "video/mp4", + ".webm": "video/webm", + ".mov": "video/quicktime", + ".json": "application/json", + ".txt": "text/plain", + } + return content_types.get(ext, "application/octet-stream") + + +def validate_file_for_discord(file_path: str) -> Tuple[bool, str]: + """ + Validate that a file is compatible with Discord uploads. + + Args: + file_path: Path to the file + + Returns: + Tuple of (is_valid, message) + """ + if not os.path.exists(file_path): + return False, f"File does not exist: {file_path}" + + file_size = os.path.getsize(file_path) + + if file_size == 0: + return False, "File is empty" + + if file_size < 1024: + return False, f"File is suspiciously small: {file_size} bytes" + + max_size = 25 * 1024 * 1024 # 25MB + if file_size > max_size: + return False, f"File exceeds Discord's size limit: {file_size} bytes (max {max_size} bytes)" + + # Check extension + ext = os.path.splitext(file_path)[1].lower().lstrip('.') + supported_formats = ['png', 'jpg', 'jpeg', 'gif', 'webp', 'mp4', 'webm', 'mov'] + + if ext in supported_formats: + return True, f"Valid {ext.upper()} file for Discord" + else: + return False, f"Format {ext} may not be fully supported by Discord" + + +def send_to_discord_with_retry( + webhook_url: str, + files: Optional[List] = None, + data: Optional[Dict] = None, + json_data: Optional[Dict] = None, + max_retries: int = 3, + timeout: int = 60 +) -> requests.Response: + """ + Send a request to Discord webhook with retry logic. + + This is a drop-in replacement for requests.post() with added retry logic, + rate limit handling, and exponential backoff. + + Args: + webhook_url: The Discord webhook URL + files: Files to upload (same format as requests.post) + data: Form data (same format as requests.post) + json_data: JSON data (same format as requests.post) + max_retries: Maximum number of retry attempts + timeout: Request timeout in seconds + + Returns: + The response object from the successful request + + Raises: + requests.exceptions.RequestException: If all retries fail + """ + last_exception = None + + for attempt in range(max_retries): + try: + if files: + response = requests.post( + webhook_url, + files=files, + data=data, + timeout=timeout + ) + elif json_data: + response = requests.post( + webhook_url, + json=json_data, + timeout=timeout + ) + else: + response = requests.post( + webhook_url, + data=data, + timeout=timeout + ) + + # Handle rate limiting + if response.status_code == 429: + retry_after = 1 + try: + retry_after = response.json().get("retry_after", 1) + except: + pass + logger.warning(f"Rate limited by Discord, waiting {retry_after}s before retry...") + time.sleep(retry_after) + continue + + # Success or client error (don't retry client errors) + if response.status_code < 500: + return response + + # Server error - retry + logger.warning(f"Discord server error {response.status_code}, attempt {attempt + 1}/{max_retries}") + + except requests.exceptions.Timeout: + logger.warning(f"Request timeout, attempt {attempt + 1}/{max_retries}") + last_exception = requests.exceptions.Timeout("Discord request timed out") + except requests.exceptions.RequestException as e: + logger.warning(f"Request error: {e}, attempt {attempt + 1}/{max_retries}") + last_exception = e + + # Exponential backoff before retry + if attempt < max_retries - 1: + wait_time = 2 ** attempt + time.sleep(wait_time) + + # If we get here, all retries failed + if last_exception: + raise last_exception + + # Return the last response even if it was an error + return response diff --git a/discordsend_utils/github_integration.py b/discordsend_utils/github_integration.py new file mode 100644 index 0000000..879216e --- /dev/null +++ b/discordsend_utils/github_integration.py @@ -0,0 +1,137 @@ +""" +GitHub Integration Utilities for ComfyUI-DiscordSend + +Handles updating GitHub repositories with Discord CDN URLs. +""" + +import base64 +import time +from typing import List, Optional, Tuple + +import requests + + +def update_github_cdn_urls( + github_repo: str, + github_token: str, + file_path: str, + cdn_urls: List[Tuple[str, str]], + commit_message: Optional[str] = None +) -> Tuple[bool, str]: + """ + Update a file in a GitHub repository with Discord CDN URLs. + + Args: + github_repo: The GitHub repository (format: username/repo) + github_token: The GitHub personal access token for authentication + file_path: The path to the file within the repository to update + cdn_urls: List of (filename, url) tuples containing Discord CDN URLs + commit_message: Optional commit message, defaults to a standard message + + Returns: + Tuple of (success, message) where success is a boolean and message is a status message + """ + # Validate required parameters + if not github_repo: + return False, "Missing GitHub repository name" + + if not github_token: + return False, "Missing GitHub personal access token" + + if not file_path: + return False, "Missing file path in repository" + + if not cdn_urls: + return False, "No CDN URLs to update" + + # Ensure repository format is valid + if "/" not in github_repo: + return False, f"Invalid GitHub repository format: {github_repo}. Expected format: username/repo" + + # Setup API endpoint + api_url = f"https://api.github.com/repos/{github_repo}/contents/{file_path}" + + headers = { + "Authorization": f"token {github_token}", + "Accept": "application/vnd.github.v3+json" + } + + try: + # Check if file exists and get its SHA + file_sha = None + current_content = "" + + response = requests.get(api_url, headers=headers, timeout=30) + + if response.status_code == 200: + file_data = response.json() + file_sha = file_data.get("sha") + + # Get current content + if file_data.get("content"): + current_content = base64.b64decode(file_data["content"]).decode("utf-8") + + elif response.status_code == 404: + pass # File doesn't exist, will create new + else: + return False, f"Error checking GitHub file: {response.status_code} - {response.text}" + + # Prepare file content + timestamp = time.strftime("%Y-%m-%d %H:%M:%S") + + # Merge with existing content if present + existing_urls = {} + if current_content: + for line in current_content.splitlines(): + if ": https://" in line and "cdn.discordapp.com" in line: + parts = line.split(": ", 1) + if len(parts) == 2: + name_part = parts[0] + # Remove numbering if present + if ". " in name_part: + name_part = name_part.split(". ", 1)[1] + existing_urls[name_part] = parts[1] + + # Add new URLs (overwrites duplicates) + for filename, url in cdn_urls: + existing_urls[filename] = url + + # Format content + new_content = f"# Discord CDN URLs\nLast updated: {timestamp}\n\n" + for i, (filename, url) in enumerate(existing_urls.items(), 1): + new_content += f"{i}. {filename}: {url}\n" + + # Set commit message + if not commit_message: + commit_message = f"Update Discord CDN URLs - {timestamp}" + + # Prepare request data + data = { + "message": commit_message, + "content": base64.b64encode(new_content.encode("utf-8")).decode("utf-8"), + } + + if file_sha: + data["sha"] = file_sha + + # Create/update file + response = requests.put(api_url, headers=headers, json=data, timeout=30) + + if response.status_code in [200, 201]: + return True, f"Successfully updated GitHub file with {len(cdn_urls)} Discord CDN URLs" + else: + return False, f"Error updating GitHub file: {response.status_code} - {response.text}" + + except requests.exceptions.Timeout: + return False, "GitHub API request timed out" + except requests.exceptions.RequestException as e: + # Scrub token from error messages + error_message = str(e) + if github_token and github_token in error_message: + error_message = error_message.replace(github_token, "[REDACTED_TOKEN]") + return False, f"GitHub API request failed: {error_message}" + except Exception as e: + error_message = str(e) + if github_token and github_token in error_message: + error_message = error_message.replace(github_token, "[REDACTED_TOKEN]") + return False, f"Exception during GitHub update: {error_message}" diff --git a/discordsend_utils/logging_config.py b/discordsend_utils/logging_config.py new file mode 100644 index 0000000..4bff9c0 --- /dev/null +++ b/discordsend_utils/logging_config.py @@ -0,0 +1,43 @@ +""" +Logging configuration for ComfyUI-DiscordSend + +Provides a configured logger for the extension. +""" + +import logging +import sys + + +def get_logger(name: str = "comfyui_discordsend") -> logging.Logger: + """ + Get a configured logger for the extension. + + Args: + name: The logger name (default: comfyui_discordsend) + + Returns: + Configured logger instance + """ + logger = logging.getLogger(name) + + # Only configure if not already configured + if not logger.handlers: + logger.setLevel(logging.DEBUG) + + # Console handler + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(logging.INFO) + + # Format + formatter = logging.Formatter( + '[%(name)s] %(levelname)s: %(message)s' + ) + console_handler.setFormatter(formatter) + + logger.addHandler(console_handler) + + return logger + + +# Default logger instance +logger = get_logger() diff --git a/discordsend_utils/prompt_extractor.py b/discordsend_utils/prompt_extractor.py new file mode 100644 index 0000000..71ffb6c --- /dev/null +++ b/discordsend_utils/prompt_extractor.py @@ -0,0 +1,240 @@ +""" +Prompt Extraction Utilities for ComfyUI-DiscordSend + +Extracts positive and negative prompts from ComfyUI workflow data. +""" + +import json +from typing import Any, Dict, List, Optional, Tuple, Union + + +# Common negative prompt indicators +NEGATIVE_INDICATORS = [ + "bad quality", "deformed", "blurry", "low quality", "worst quality", + "ugly", "disfigured", "low res", "poorly drawn", "mutation", + "extra limbs", "bad anatomy", "watermark", "text", "signature" +] + + +def extract_prompts_from_workflow(workflow_data: Any) -> Tuple[Optional[str], Optional[str]]: + """ + Extract positive and negative prompts from workflow data. + + Analyzes ComfyUI workflow structure to find CLIPTextEncode nodes and + determine which contains the positive vs negative prompt. + + Args: + workflow_data: The workflow data dictionary or JSON string + + Returns: + A tuple of (positive_prompt, negative_prompt) or (None, None) if not found + """ + if workflow_data is None: + return None, None + + # Convert string to dict if necessary + if isinstance(workflow_data, str): + try: + data = json.loads(workflow_data) + except json.JSONDecodeError: + return None, None + else: + data = workflow_data + + if not isinstance(data, dict): + return None, None + + positive_prompt = None + negative_prompt = None + + # Find CLIPTextEncode nodes + if "nodes" not in data: + return None, None + + nodes = data["nodes"] + clip_nodes = _find_clip_text_encode_nodes(nodes) + + if not clip_nodes: + return None, None + + # Determine positive/negative based on content and structure + if len(clip_nodes) == 1: + # Single CLIP node - assume it's the positive prompt + positive_prompt = _get_prompt_text(clip_nodes[0]) + elif len(clip_nodes) >= 2: + # Multiple CLIP nodes - need to determine which is which + positive_prompt, negative_prompt = _classify_prompts(clip_nodes, data) + + # Ensure we return empty string for negative if we have positive but not negative + if positive_prompt is not None and negative_prompt is None: + negative_prompt = "" + + return positive_prompt, negative_prompt + + +def _find_clip_text_encode_nodes(nodes: Union[List, Dict]) -> List[Dict]: + """Find all CLIPTextEncode nodes in the workflow.""" + clip_nodes = [] + + if isinstance(nodes, list): + for node in nodes: + if _is_clip_text_encode(node): + clip_nodes.append(node) + elif isinstance(nodes, dict): + for node_id, node in nodes.items(): + if _is_clip_text_encode(node): + node_copy = dict(node) + node_copy["id"] = node_id + clip_nodes.append(node_copy) + + return clip_nodes + + +def _is_clip_text_encode(node: Any) -> bool: + """Check if a node is a CLIPTextEncode node with valid text.""" + if not isinstance(node, dict): + return False + + if node.get("type") != "CLIPTextEncode": + return False + + widgets = node.get("widgets_values", []) + if not isinstance(widgets, list) or len(widgets) == 0: + return False + + # Must have a string value at position 0 + if not isinstance(widgets[0], str): + return False + + return True + + +def _get_prompt_text(node: Dict) -> Optional[str]: + """Extract the prompt text from a CLIP node.""" + widgets = node.get("widgets_values", []) + if widgets and isinstance(widgets[0], str): + return widgets[0] + return None + + +def _classify_prompts(clip_nodes: List[Dict], workflow_data: Dict) -> Tuple[Optional[str], Optional[str]]: + """ + Classify which CLIP nodes contain positive vs negative prompts. + + Uses multiple heuristics: + 1. Content analysis (negative prompts often contain quality-related terms) + 2. Connection analysis (traces connections to sampler nodes) + """ + positive_prompt = None + negative_prompt = None + + # First pass: content-based classification + for node in clip_nodes: + prompt_text = _get_prompt_text(node) + if prompt_text is None: + continue + + prompt_lower = prompt_text.lower() + + # Count negative indicators + negative_matches = sum(1 for indicator in NEGATIVE_INDICATORS if indicator in prompt_lower) + + if negative_matches >= 3: + # Likely a negative prompt + if negative_prompt is None: + negative_prompt = prompt_text + else: + # Assume positive + if positive_prompt is None: + positive_prompt = prompt_text + + # Second pass: connection-based classification (if needed) + if positive_prompt is None or negative_prompt is None: + positive_prompt, negative_prompt = _classify_by_connections( + clip_nodes, workflow_data, positive_prompt, negative_prompt + ) + + # Fallback: if we still can't determine, use first two nodes + if positive_prompt is None and negative_prompt is None and len(clip_nodes) >= 2: + # Convention: assume first is positive, second is negative + positive_prompt = _get_prompt_text(clip_nodes[0]) + negative_prompt = _get_prompt_text(clip_nodes[1]) + elif positive_prompt is None and negative_prompt is not None: + # Find the other prompt + for node in clip_nodes: + text = _get_prompt_text(node) + if text != negative_prompt: + positive_prompt = text + break + elif negative_prompt is None and positive_prompt is not None: + # Find the other prompt + for node in clip_nodes: + text = _get_prompt_text(node) + if text != positive_prompt: + negative_prompt = text + break + + return positive_prompt, negative_prompt + + +def _classify_by_connections( + clip_nodes: List[Dict], + workflow_data: Dict, + existing_positive: Optional[str], + existing_negative: Optional[str] +) -> Tuple[Optional[str], Optional[str]]: + """Try to classify prompts by analyzing node connections to samplers.""" + positive = existing_positive + negative = existing_negative + + links = workflow_data.get("links", []) + if not isinstance(links, list): + return positive, negative + + nodes = workflow_data.get("nodes", []) + + # Find sampler nodes + samplers = [] + if isinstance(nodes, list): + for node in nodes: + if isinstance(node, dict) and "KSampler" in node.get("type", ""): + samplers.append(node) + elif isinstance(nodes, dict): + for node_id, node in nodes.items(): + if isinstance(node, dict) and "KSampler" in node.get("type", ""): + node_copy = dict(node) + node_copy["id"] = node_id + samplers.append(node_copy) + + if not samplers: + return positive, negative + + # Trace connections from CLIP nodes to samplers + for link in links: + if len(link) < 4: + continue + + from_node_id = link[0] + to_node_id = link[2] + to_slot = link[3] + + # Find matching CLIP node and sampler + for clip_node in clip_nodes: + clip_id = clip_node.get("id") + if clip_id == from_node_id: + for sampler in samplers: + if sampler.get("id") == to_node_id: + # Determine slot type + inputs = sampler.get("inputs", {}) + if isinstance(inputs, dict): + input_keys = list(inputs.keys()) + if to_slot < len(input_keys): + input_name = input_keys[to_slot].lower() + prompt_text = _get_prompt_text(clip_node) + + if "positive" in input_name and positive is None: + positive = prompt_text + elif "negative" in input_name and negative is None: + negative = prompt_text + + return positive, negative diff --git a/discordsend_utils/sanitizer.py b/discordsend_utils/sanitizer.py new file mode 100644 index 0000000..ea2d0d0 --- /dev/null +++ b/discordsend_utils/sanitizer.py @@ -0,0 +1,269 @@ +""" +JSON Sanitization Utilities for ComfyUI-DiscordSend + +Provides comprehensive sanitization of sensitive data (webhook URLs, GitHub tokens) +from JSON data structures before export or sharing. +""" + +import json +import re +from typing import Any, Dict, List, Optional, Union + + +# Patterns for detecting sensitive data +WEBHOOK_PATTERNS = [ + r"discord\.com/api/webhooks", + r"discordapp\.com/api/webhooks", +] + +GITHUB_TOKEN_PREFIXES = [ + "ghp_", # GitHub personal access token + "github_pat_", # GitHub personal access token (new format) + "gho_", # GitHub OAuth token + "ghs_", # GitHub service token + "ghu_", # GitHub user-to-server token +] + + +def is_webhook_url(value: str) -> bool: + """Check if a string appears to be a Discord webhook URL.""" + if not isinstance(value, str): + return False + + for pattern in WEBHOOK_PATTERNS: + if re.search(pattern, value, re.IGNORECASE): + return True + + # Also check for generic webhook patterns in URLs + if value.startswith("http") and ("webhook" in value.lower() or "discord" in value.lower()): + return True + + return False + + +def is_github_token(value: str) -> bool: + """Check if a string appears to be a GitHub token.""" + if not isinstance(value, str): + return False + + for prefix in GITHUB_TOKEN_PREFIXES: + if value.startswith(prefix): + return True + + return False + + +def is_potential_token(value: str, context_type: str = "") -> bool: + """Check if a string might be a token based on length and context.""" + if not isinstance(value, str): + return False + + # Long alphanumeric strings in GitHub-related contexts + if len(value) >= 40 and "github" in context_type.lower(): + if any(c.isalnum() for c in value): + return True + + return False + + +def sanitize_string(value: str, context_type: str = "") -> str: + """ + Sanitize a single string value. + + Args: + value: The string to sanitize + context_type: Optional context (e.g., node type) to help identify tokens + + Returns: + Empty string if sensitive data detected, original value otherwise + """ + if is_webhook_url(value): + return "" + + if is_github_token(value): + return "" + + if is_potential_token(value, context_type): + return "" + + return value + + +def sanitize_widget_values(widgets: List, node_type: str = "") -> List: + """ + Sanitize a list of widget values from a ComfyUI node. + + Args: + widgets: List of widget values + node_type: The type of node for context + + Returns: + Sanitized list with sensitive values replaced with empty strings + """ + result = [] + for value in widgets: + if isinstance(value, str): + result.append(sanitize_string(value, node_type)) + elif isinstance(value, dict): + result.append(sanitize_dict(value)) + elif isinstance(value, list): + result.append(sanitize_widget_values(value, node_type)) + else: + result.append(value) + return result + + +def sanitize_node_inputs(inputs: Dict, node_type: str = "") -> Dict: + """ + Sanitize the inputs dictionary of a ComfyUI node. + + Args: + inputs: Dictionary of node inputs + node_type: The type of node for context + + Returns: + Sanitized dictionary + """ + result = {} + for key, value in inputs.items(): + if key in ("webhook_url", "github_token"): + result[key] = "" + elif isinstance(value, str): + result[key] = sanitize_string(value, node_type) + elif isinstance(value, dict): + result[key] = sanitize_node_inputs(value, node_type) + elif isinstance(value, list): + result[key] = sanitize_widget_values(value, node_type) + else: + result[key] = value + return result + + +def sanitize_node(node: Dict) -> Dict: + """ + Sanitize a single ComfyUI node. + + Args: + node: The node dictionary + + Returns: + Sanitized node dictionary + """ + if not isinstance(node, dict): + return node + + result = dict(node) + node_type = result.get("type", "") + + # Sanitize inputs + if "inputs" in result and isinstance(result["inputs"], dict): + result["inputs"] = sanitize_node_inputs(result["inputs"], node_type) + + # Sanitize widget values + if "widgets_values" in result and isinstance(result["widgets_values"], list): + result["widgets_values"] = sanitize_widget_values(result["widgets_values"], node_type) + + return result + + +def sanitize_dict(data: Dict) -> Dict: + """ + Recursively sanitize a dictionary. + + Args: + data: Dictionary to sanitize + + Returns: + Sanitized dictionary + """ + result = {} + + for key, value in data.items(): + # Handle known sensitive keys + if key in ("webhook_url", "github_token"): + result[key] = "" + continue + + # Handle nested structures + if isinstance(value, dict): + result[key] = sanitize_dict(value) + elif isinstance(value, list): + result[key] = sanitize_list(value) + elif isinstance(value, str): + result[key] = sanitize_string(value) + else: + result[key] = value + + # Special handling for ComfyUI workflow structure + if "nodes" in result: + nodes = result["nodes"] + if isinstance(nodes, list): + result["nodes"] = [sanitize_node(n) for n in nodes] + elif isinstance(nodes, dict): + result["nodes"] = {k: sanitize_node(v) for k, v in nodes.items()} + + return result + + +def sanitize_list(data: List) -> List: + """ + Recursively sanitize a list. + + Args: + data: List to sanitize + + Returns: + Sanitized list + """ + result = [] + for item in data: + if isinstance(item, dict): + result.append(sanitize_dict(item)) + elif isinstance(item, list): + result.append(sanitize_list(item)) + elif isinstance(item, str): + result.append(sanitize_string(item)) + else: + result.append(item) + return result + + +def sanitize_json_for_export(json_data: Any) -> Any: + """ + Remove sensitive webhook data and GitHub tokens from JSON data to protect user security. + + This is the main entry point for sanitization. It handles: + - Discord webhook URLs (multiple formats) + - GitHub personal access tokens (all variants) + - Nested dictionaries and lists + - ComfyUI workflow structures + + Args: + json_data: The JSON data object (dict, list, or string) to sanitize + + Returns: + The sanitized JSON data with sensitive information removed + """ + if json_data is None: + return None + + # Handle string input (may be JSON string or plain string) + if isinstance(json_data, str): + try: + data = json.loads(json_data) + result = sanitize_json_for_export(data) + return json.dumps(result) + except json.JSONDecodeError: + # Not valid JSON - check if it's sensitive data directly + return sanitize_string(json_data) + + # Handle dictionary + if isinstance(json_data, dict): + return sanitize_dict(json_data) + + # Handle list + if isinstance(json_data, list): + return sanitize_list(json_data) + + # Return other types as-is + return json_data diff --git a/discordsend_utils/workflow_builder.py b/discordsend_utils/workflow_builder.py new file mode 100644 index 0000000..09f26d9 --- /dev/null +++ b/discordsend_utils/workflow_builder.py @@ -0,0 +1,146 @@ +import json +import random +import logging +from typing import Dict, Any, Optional, Tuple, List + +logger = logging.getLogger(__name__) + +class WorkflowBuilder: + """Helper to manipulate ComfyUI workflow JSONs.""" + + def __init__(self, workflow_json: Dict[str, Any]): + self.workflow = workflow_json + # Create lookups + self._nodes = self.workflow + if "nodes" in self.workflow and isinstance(self.workflow["nodes"], list): + # Handle "graph" format vs "api" format if needed. + # But usually for API we stick to the {node_id: node_data} format. + # If input is graph format, it might need conversion or distinct handling. + # Assuming API format for now as that's what's sent to /prompt. + pass + + @classmethod + def from_json_string(cls, json_str: str) -> 'WorkflowBuilder': + """Load from JSON int.""" + return cls(json.loads(json_str)) + + def get_workflow(self) -> Dict[str, Any]: + """Get the current workflow dict.""" + return self.workflow + + def set_prompt(self, positive: str, negative: Optional[str] = None) -> None: + """ + Attempt to set positive and negative prompts. + Heuristics: + - Look for CLIPTextEncode nodes. + - Often one is connected to KSampler 'positive' and one to 'negative'. + - Or look for custom titles like 'Positive Prompt', 'Negative Prompt'. + """ + # Simple heuristic: Find CLIPTextEncode nodes + # If we have title/coloring, we can use that. + # Otherwise, we might need graph traversal to see what connects to KSampler. + + # For MVP, let's assume standard ComfyUI structure or look for specific titles first + + positive_node_id = self._find_node_by_title("Positive Prompt") + negative_node_id = self._find_node_by_title("Negative Prompt") + + # Fallback: Find KSampler and trace back + if not positive_node_id or not negative_node_id: + ksampler_id, ksampler = self._find_node_by_class("KSampler") + if ksampler: + # KSampler inputs: model, positive, negative, latent_image + if not positive_node_id: + positive_node_id = self._trace_input(ksampler, "positive") + if not negative_node_id: + negative_node_id = self._trace_input(ksampler, "negative") + + if positive_node_id: + self._update_node_input(positive_node_id, "text", positive) + else: + logger.warning("Could not identify Positive Prompt node.") + + if negative and negative_node_id: + self._update_node_input(negative_node_id, "text", negative) + elif negative: + logger.warning("Could not identify Negative Prompt node.") + + def set_seed(self, seed: int) -> int: + """Set seed on KSampler nodes or Seed nodes.""" + # Find KSampler or anything with a 'seed' widget + updated = False + for node_id, node in self.workflow.items(): + if "inputs" in node: + if "seed" in node["inputs"]: + # Ensure it's an int widget, not a link + if isinstance(node["inputs"]["seed"], (int, float)) or (isinstance(node["inputs"]["seed"], str) and node["inputs"]["seed"].isdigit()): + node["inputs"]["seed"] = seed + updated = True + if "noise_seed" in node["inputs"]: + # Some nodes call it noise_seed + if isinstance(node["inputs"]["noise_seed"], (int, float)): + node["inputs"]["noise_seed"] = seed + updated = True + + if not updated: + logger.warning("Could not find any seed inputs to update.") + return seed + + def set_image_dimensions(self, width: int, height: int) -> None: + """Set width and height on EmptyLatentImage nodes.""" + node_id, _ = self._find_node_by_class("EmptyLatentImage") + if node_id: + self._update_node_input(node_id, "width", width) + self._update_node_input(node_id, "height", height) + + def set_steps(self, steps: int) -> None: + """Set steps on KSampler.""" + ksampler_ids = self._find_nodes_by_class("KSampler") + for nid in ksampler_ids: + self._update_node_input(nid, "steps", steps) + + def set_cfg(self, cfg: float) -> None: + """Set CFG scale on KSampler.""" + ksampler_ids = self._find_nodes_by_class("KSampler") + for nid in ksampler_ids: + self._update_node_input(nid, "cfg", cfg) + + def _find_node_by_title(self, title: str) -> Optional[str]: + """Find node by its custom title (`_meta.title`).""" + for node_id, node in self.workflow.items(): + if "_meta" in node and node["_meta"].get("title") == title: + return node_id + return None + + def _find_node_by_class(self, class_type: str) -> Tuple[Optional[str], Optional[Dict]]: + """Find first node of a specific class type.""" + for node_id, node in self.workflow.items(): + if node.get("class_type") == class_type: + return node_id, node + return None, None + + def _find_nodes_by_class(self, class_type: str) -> List[str]: + """Find all nodes of a specific class type.""" + ids = [] + for node_id, node in self.workflow.items(): + if node.get("class_type") == class_type: + ids.append(node_id) + return ids + + def _trace_input(self, node: Dict, input_name: str) -> Optional[str]: + """ + Trace back an input link to find the source node. + Input format in API JSON: "input_name": ["source_node_id", slot_index] + """ + if "inputs" not in node or input_name not in node["inputs"]: + return None + + link = node["inputs"][input_name] + # Link structure: [node_id, slot_idx] + if isinstance(link, list) and len(link) == 2: + return str(link[0]) + return None + + def _update_node_input(self, node_id: str, input_name: str, value: Any) -> None: + if node_id in self.workflow and "inputs" in self.workflow[node_id]: + self.workflow[node_id]["inputs"][input_name] = value diff --git a/scripts/auto_review.sh b/scripts/auto_review.sh new file mode 100755 index 0000000..d1694bf --- /dev/null +++ b/scripts/auto_review.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +# auto_review.sh - Automatically review the latest changes using Gemini CLI + +# 1. Get the latest changes (staged or last commit) +# If there are staged changes, review them. Otherwise review last commit. +if git diff --quiet --cached; then + # No staged changes, check last commit + DIFF_CONTENT=$(git show HEAD) + CONTEXT="Review the following changes from the last commit:" +else + # Staged changes exist + DIFF_CONTENT=$(git diff --cached) + CONTEXT="Review the following staged changes:" +fi + +if [ -z "$DIFF_CONTENT" ]; then + echo "No changes found to review." + exit 0 +fi + +# 2. Construct Prompt +PROMPT="You are a Senior Software Engineer acting as a code reviewer. +Review the following code changes for: +1. Potential bugs or race conditions +2. Security vulnerabilities +3. Code style and best practices +4. Logical errors + +Be concise and constructive. + +$CONTEXT +\`\`\`diff +$DIFF_CONTENT +\`\`\` +" + +# 3. Call Gemini CLI +echo "🤖 Asking Gemini to review changes..." +echo "----------------------------------------" +gemini "$PROMPT" +echo "----------------------------------------" diff --git a/tests/test_utils.py b/tests/test_utils.py index be674c8..0673707 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -12,8 +12,8 @@ import unittest # Add parent directory to path for imports sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from utils.sanitizer import sanitize_json_for_export -from utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging +from discordsend_utils.sanitizer import sanitize_json_for_export +from discordsend_utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging class TestSanitizer(unittest.TestCase):