refactor(core): restructure project into shared/ and nodes/ modules
This commit is contained in:
@@ -0,0 +1,40 @@
|
||||
"""
|
||||
ComfyUI-DiscordSend Shared Utilities
|
||||
|
||||
This package contains shared utilities used by both ComfyUI nodes and the Discord bot.
|
||||
Organized into subpackages:
|
||||
- discord: Discord webhook and message utilities
|
||||
- media: Image and video processing utilities
|
||||
- workflow: ComfyUI workflow manipulation utilities
|
||||
"""
|
||||
|
||||
# Re-export commonly used utilities for convenience
|
||||
from .workflow.sanitizer import sanitize_json_for_export
|
||||
from .workflow.prompt_extractor import extract_prompts_from_workflow
|
||||
from .workflow.workflow_builder import WorkflowBuilder
|
||||
from .discord.webhook_client import (
|
||||
DiscordWebhookClient,
|
||||
validate_webhook_url,
|
||||
send_to_discord_with_retry
|
||||
)
|
||||
from .media.image_processing import tensor_to_numpy_uint8
|
||||
from .github_integration import update_github_cdn_urls
|
||||
from .logging_config import setup_logging, get_logger
|
||||
|
||||
__all__ = [
|
||||
# Workflow utilities
|
||||
'sanitize_json_for_export',
|
||||
'extract_prompts_from_workflow',
|
||||
'WorkflowBuilder',
|
||||
# Discord utilities
|
||||
'DiscordWebhookClient',
|
||||
'validate_webhook_url',
|
||||
'send_to_discord_with_retry',
|
||||
# Media utilities
|
||||
'tensor_to_numpy_uint8',
|
||||
# GitHub integration
|
||||
'update_github_cdn_urls',
|
||||
# Logging
|
||||
'setup_logging',
|
||||
'get_logger',
|
||||
]
|
||||
@@ -0,0 +1,21 @@
|
||||
"""
|
||||
Discord Integration Utilities
|
||||
|
||||
Provides webhook client, message building, and CDN URL handling.
|
||||
"""
|
||||
|
||||
from .webhook_client import (
|
||||
DiscordWebhookClient,
|
||||
validate_webhook_url,
|
||||
sanitize_webhook_for_logging,
|
||||
send_to_discord_with_retry,
|
||||
validate_file_for_discord
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'DiscordWebhookClient',
|
||||
'validate_webhook_url',
|
||||
'sanitize_webhook_for_logging',
|
||||
'send_to_discord_with_retry',
|
||||
'validate_file_for_discord',
|
||||
]
|
||||
@@ -0,0 +1,444 @@
|
||||
"""
|
||||
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 json
|
||||
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("https://"):
|
||||
return False, "Webhook URL must start with https://"
|
||||
|
||||
# Check against known patterns
|
||||
for pattern in WEBHOOK_URL_PATTERNS:
|
||||
if re.match(pattern, url, re.IGNORECASE):
|
||||
return True, "Valid 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": json.dumps(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:
|
||||
# Sanitize error message to prevent token leakage
|
||||
error_msg = str(e)
|
||||
# Use case-insensitive matching to handle uppercase URLs
|
||||
match = re.search(r"/api/webhooks/\d+/([\w-]+)", self.webhook_url, re.IGNORECASE)
|
||||
if match:
|
||||
token = match.group(1)
|
||||
if token in error_msg:
|
||||
error_msg = error_msg.replace(token, "[REDACTED]")
|
||||
|
||||
last_error = error_msg
|
||||
|
||||
# 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
|
||||
ValueError: If the webhook URL is invalid
|
||||
"""
|
||||
# Validate URL to prevent SSRF
|
||||
is_valid, error_msg = validate_webhook_url(webhook_url)
|
||||
if not is_valid:
|
||||
raise ValueError(f"Invalid webhook URL: {error_msg}")
|
||||
|
||||
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:
|
||||
# Sanitize error message to prevent token leakage
|
||||
error_msg = str(e)
|
||||
# Use case-insensitive matching to handle uppercase URLs
|
||||
match = re.search(r"/api/webhooks/\d+/([\w-]+)", webhook_url, re.IGNORECASE)
|
||||
if match:
|
||||
token = match.group(1)
|
||||
if token in error_msg:
|
||||
error_msg = error_msg.replace(token, "[REDACTED]")
|
||||
|
||||
logger.warning(f"Request error: {error_msg}, attempt {attempt + 1}/{max_retries}")
|
||||
|
||||
# Store sanitized exception to avoid leaking token if raised later
|
||||
if match and match.group(1) in str(e):
|
||||
# Create a new exception of the same type with sanitized message
|
||||
# We try to preserve the exception type, but fallback to RequestException if init fails
|
||||
try:
|
||||
last_exception = type(e)(error_msg)
|
||||
# Preserve context attributes if possible
|
||||
last_exception.request = getattr(e, "request", None)
|
||||
last_exception.response = getattr(e, "response", None)
|
||||
except:
|
||||
last_exception = requests.exceptions.RequestException(error_msg)
|
||||
last_exception.request = getattr(e, "request", None)
|
||||
last_exception.response = getattr(e, "response", None)
|
||||
else:
|
||||
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
|
||||
@@ -0,0 +1,181 @@
|
||||
"""
|
||||
GitHub Integration Utilities for ComfyUI-DiscordSend
|
||||
|
||||
Handles updating GitHub repositories with Discord CDN URLs.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import time
|
||||
import re
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def validate_github_repo(repo: str) -> bool:
|
||||
"""
|
||||
Validate GitHub repository format (username/repo).
|
||||
Strictly enforces alphanumeric, hyphens, underscores, and periods.
|
||||
Prevents path traversal and injection.
|
||||
"""
|
||||
if not repo:
|
||||
return False
|
||||
# Pattern: username/repo
|
||||
# GitHub usernames: alphanumeric, hyphens (max 39 chars)
|
||||
# Repo names: alphanumeric, hyphens, periods, underscores
|
||||
pattern = r"^[a-zA-Z0-9-]+/[\w.-]+$"
|
||||
return bool(re.match(pattern, repo))
|
||||
|
||||
|
||||
def validate_file_path(path: str) -> bool:
|
||||
"""
|
||||
Validate file path for GitHub API.
|
||||
Prevents path traversal (..) and absolute paths.
|
||||
"""
|
||||
if not path:
|
||||
return False
|
||||
# Prevent traversal
|
||||
if ".." in path:
|
||||
return False
|
||||
# Prevent absolute paths (GitHub API treats paths as relative to root)
|
||||
if path.startswith("/"):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
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"
|
||||
|
||||
# Strictly validate repository format to prevent traversal/injection
|
||||
if not validate_github_repo(github_repo):
|
||||
return False, f"Invalid GitHub repository format: {github_repo}. Expected format: username/repo (alphanumeric, hyphens, periods, underscores only)"
|
||||
|
||||
# Strictly validate file path to prevent traversal
|
||||
if not validate_file_path(file_path):
|
||||
return False, f"Invalid file path: {file_path}. Path traversal (..) and absolute paths are not allowed."
|
||||
|
||||
# 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:
|
||||
# Sanitize response text
|
||||
error_details = response.text
|
||||
if github_token and github_token in error_details:
|
||||
error_details = error_details.replace(github_token, "[REDACTED_TOKEN]")
|
||||
return False, f"Error checking GitHub file: {response.status_code} - {error_details}"
|
||||
|
||||
# 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:
|
||||
# Sanitize response text to ensure no token leakage
|
||||
error_details = response.text
|
||||
if github_token and github_token in error_details:
|
||||
error_details = error_details.replace(github_token, "[REDACTED_TOKEN]")
|
||||
return False, f"Error updating GitHub file: {response.status_code} - {error_details}"
|
||||
|
||||
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}"
|
||||
@@ -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()
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
Media Processing Utilities
|
||||
|
||||
Provides image and video processing functions.
|
||||
"""
|
||||
|
||||
from .image_processing import tensor_to_numpy_uint8
|
||||
|
||||
__all__ = [
|
||||
'tensor_to_numpy_uint8',
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
"""
|
||||
Image processing utilities for ComfyUI-DiscordSend.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
def tensor_to_numpy_uint8(tensor: torch.Tensor) -> np.ndarray:
|
||||
"""
|
||||
Convert a PyTorch tensor (0-1 float) to a numpy uint8 array (0-255).
|
||||
|
||||
This function optimizes performance by doing scaling, clamping, and casting
|
||||
in PyTorch before moving data to CPU/NumPy, avoiding large intermediate float arrays.
|
||||
|
||||
Args:
|
||||
tensor: PyTorch tensor with values in range [0, 1]
|
||||
|
||||
Returns:
|
||||
Numpy uint8 array with values in range [0, 255]
|
||||
"""
|
||||
# Optimization: Use torch operations for scaling/clipping/casting to avoid large float64 intermediate arrays on CPU
|
||||
# This is ~70% faster than naive numpy conversion: np.clip(255. * tensor.cpu().numpy(), 0, 255).astype(np.uint8)
|
||||
return (tensor * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
@@ -0,0 +1,15 @@
|
||||
"""
|
||||
Workflow Manipulation Utilities
|
||||
|
||||
Provides sanitization, prompt extraction, and workflow building tools.
|
||||
"""
|
||||
|
||||
from .sanitizer import sanitize_json_for_export
|
||||
from .prompt_extractor import extract_prompts_from_workflow
|
||||
from .workflow_builder import WorkflowBuilder
|
||||
|
||||
__all__ = [
|
||||
'sanitize_json_for_export',
|
||||
'extract_prompts_from_workflow',
|
||||
'WorkflowBuilder',
|
||||
]
|
||||
@@ -0,0 +1,256 @@
|
||||
"""
|
||||
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" in data:
|
||||
nodes = data["nodes"]
|
||||
else:
|
||||
# Check if it's API format (dict of nodes)
|
||||
nodes = data
|
||||
|
||||
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
|
||||
|
||||
# Handle both Workflow format (type) and API format (class_type)
|
||||
node_type = node.get("type") or node.get("class_type")
|
||||
if node_type != "CLIPTextEncode":
|
||||
return False
|
||||
|
||||
# Check for text in either widgets_values (Workflow) or inputs (API)
|
||||
text = _get_prompt_text(node)
|
||||
return text is not None
|
||||
|
||||
|
||||
def _get_prompt_text(node: Dict) -> Optional[str]:
|
||||
"""Extract the prompt text from a CLIP node."""
|
||||
# Workflow format (widgets_values)
|
||||
widgets = node.get("widgets_values", [])
|
||||
if isinstance(widgets, list) and len(widgets) > 0 and isinstance(widgets[0], str):
|
||||
return widgets[0]
|
||||
|
||||
# API format (inputs.text)
|
||||
inputs = node.get("inputs", {})
|
||||
if isinstance(inputs, dict) and isinstance(inputs.get("text"), str):
|
||||
return inputs.get("text")
|
||||
|
||||
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)
|
||||
"""
|
||||
if not clip_nodes:
|
||||
return None, None
|
||||
|
||||
# First pass: Score all nodes based on content
|
||||
node_scores = []
|
||||
for node in clip_nodes:
|
||||
prompt_text = _get_prompt_text(node)
|
||||
# Skip empty or None text
|
||||
if not prompt_text or not prompt_text.strip():
|
||||
continue
|
||||
|
||||
prompt_lower = prompt_text.lower()
|
||||
score = sum(1 for indicator in NEGATIVE_INDICATORS if indicator in prompt_lower)
|
||||
node_scores.append((node, score, prompt_text))
|
||||
|
||||
if not node_scores:
|
||||
return None, None
|
||||
|
||||
# Sort by score descending (highest score first for negative)
|
||||
node_scores.sort(key=lambda x: x[1], reverse=True)
|
||||
|
||||
positive_prompt = None
|
||||
negative_prompt = None
|
||||
|
||||
# Highest score is likely negative (if score > 0)
|
||||
if node_scores[0][1] > 0:
|
||||
negative_prompt = node_scores[0][2]
|
||||
# The node with the lowest score is likely the positive prompt
|
||||
# Since it's sorted descending, pick the last one
|
||||
positive_prompt = node_scores[-1][2]
|
||||
|
||||
# If we have more than 2 nodes and the last one is NOT the only zero-score node,
|
||||
# we might need more logic, but for 2 nodes, this is correct.
|
||||
# If positive and negative are the same (only 1 node with score > 0), clear positive.
|
||||
if positive_prompt == negative_prompt:
|
||||
positive_prompt = None
|
||||
else:
|
||||
# All scores are 0, use connection analysis
|
||||
positive_prompt, negative_prompt = _classify_by_connections(
|
||||
clip_nodes, workflow_data, None, None
|
||||
)
|
||||
|
||||
# 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 _, _, text in node_scores:
|
||||
if text != negative_prompt:
|
||||
positive_prompt = text
|
||||
break
|
||||
elif negative_prompt is None and positive_prompt is not None:
|
||||
# Find the other prompt
|
||||
for _, _, text in node_scores:
|
||||
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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user