refactor: Rename utils to discordsend_utils
This commit is contained in:
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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',
|
||||
]
|
||||
@@ -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
|
||||
@@ -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}"
|
||||
@@ -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,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
|
||||
@@ -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
|
||||
Executable
+42
@@ -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 "----------------------------------------"
|
||||
+2
-2
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user