Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
325e2610a4 | ||
|
|
e10b421b23 | ||
|
|
3e98047524 | ||
|
|
bf35dc72c6 | ||
|
|
bb8166cc8f | ||
|
|
9ca5218ce0 | ||
|
|
c498711f10 | ||
|
|
6de50104f1 | ||
|
|
348599a139 | ||
|
|
52004b753d | ||
|
|
519e6025a8 | ||
|
|
2c4277711a | ||
|
|
195842c634 | ||
|
|
851b2ef353 | ||
|
|
0b723ce1df | ||
|
|
29d2c9856a | ||
|
|
9520576e4e | ||
|
|
13d2114df9 | ||
|
|
fc4a3d37cd | ||
|
|
37a444993e | ||
|
|
128e19a9b5 | ||
|
|
d15f6ffc8c | ||
|
|
510279d4d7 | ||
|
|
1697364e13 | ||
|
|
ed2d1388ec | ||
|
|
de5d753cdf | ||
|
|
523fe68f6b | ||
|
|
ad360061d7 | ||
|
|
98bbda1f94 | ||
|
|
c4a96e3484 | ||
|
|
0e655a4303 | ||
|
|
ec09d32cf6 | ||
|
|
9c8e73e578 | ||
|
|
af9936a706 | ||
|
|
9716a63591 | ||
|
|
7775d0b212 | ||
|
|
59c04a6b84 | ||
|
|
6c4c2ebc94 | ||
|
|
89def34b23 | ||
|
|
411336a217 | ||
|
|
8a2793e3fa | ||
|
|
194654ff1a | ||
|
|
fa58cddf73 | ||
|
|
4b2a99e5f5 | ||
|
|
766168ec1b | ||
|
|
87112b8b2a | ||
|
|
73b2bfe5b1 | ||
|
|
b3d8f9d16f | ||
|
|
b63dac6cad | ||
|
|
3ee819d042 | ||
|
|
1ba7122694 | ||
|
|
e00194e0e7 | ||
|
|
c74b572c38 | ||
|
|
5d23e30dbc | ||
|
|
7382295a76 | ||
|
|
c4ef38cff3 | ||
|
|
c1b1e2497c | ||
|
|
7cf1b45545 | ||
|
|
0c7c1d83be | ||
|
|
f53010ceef | ||
|
|
8206c475e5 | ||
|
|
f2e1779450 | ||
|
|
8ffb72bb4d | ||
|
|
6af9e26255 | ||
|
|
e7fde9951a | ||
|
|
def123c242 | ||
|
|
436b773bd8 | ||
|
|
b083702eb9 | ||
|
|
618b9ae4bb | ||
|
|
569f76f484 | ||
|
|
0d6ed9d9c5 | ||
|
|
cc07d20f86 | ||
|
|
36c5d06cf8 | ||
|
|
b8aadd7faa | ||
|
|
1e94a69e5f | ||
|
|
87c97e04a3 | ||
|
|
76da850b3d | ||
|
|
a208cd482b | ||
|
|
9295546c17 | ||
|
|
852840dcae | ||
|
|
06d7b955f5 | ||
|
|
64179a3170 | ||
|
|
1a84f62ec1 | ||
|
|
d7b20fc717 | ||
|
|
95e00c65bf | ||
|
|
4d34e27463 | ||
|
|
5c6f15469d | ||
|
|
9acae29f27 | ||
|
|
83f0059720 | ||
|
|
81010bbb4e | ||
|
|
c4dcfd586d | ||
|
|
e3e95ab197 | ||
|
|
7cc72606a5 | ||
|
|
8be6270aef | ||
|
|
b944db244e | ||
|
|
d99aa60623 | ||
|
|
214ecf35eb | ||
|
|
711aeed35d | ||
|
|
a9ff597d5c | ||
|
|
e1f9b50c3d | ||
|
|
87a0a5d5fe | ||
|
|
45a1627992 | ||
|
|
1c9ab10e16 | ||
|
|
f67a978c5d | ||
|
|
291fa65a7b | ||
|
|
c5e79124cc | ||
|
|
5610f51a71 | ||
|
|
3fd583c483 | ||
|
|
8888c2cb0b | ||
|
|
35178b92c8 | ||
|
|
d35caae023 | ||
|
|
3f553e600a | ||
|
|
cd2052757a | ||
|
|
73cf5fc0dd | ||
|
|
2a37c1b410 | ||
|
|
24d6912396 | ||
|
|
15883ebb66 | ||
|
|
decd7def28 | ||
|
|
491e4a2758 | ||
|
|
a81ae85d1e | ||
|
|
f49f65a196 | ||
|
|
6558997050 | ||
|
|
cc38033f6d |
@@ -0,0 +1,48 @@
|
||||
# ComfyUI-DiscordSend Bot Environment Variables
|
||||
# Copy this file to .env and fill in your values
|
||||
# These override settings in config.yaml
|
||||
|
||||
# =============================================================================
|
||||
# REQUIRED
|
||||
# =============================================================================
|
||||
|
||||
# Discord bot token from Discord Developer Portal
|
||||
DISCORDBOT_DISCORD_TOKEN=your_bot_token_here
|
||||
|
||||
# =============================================================================
|
||||
# OPTIONAL - ComfyUI Connection
|
||||
# =============================================================================
|
||||
|
||||
# ComfyUI server URL (default: http://127.0.0.1:8188)
|
||||
# DISCORDBOT_COMFYUI_URL=http://127.0.0.1:8188
|
||||
|
||||
# ComfyUI WebSocket URL (default: auto-derived from COMFYUI_URL)
|
||||
# DISCORDBOT_COMFYUI_WS_URL=ws://127.0.0.1:8188/ws
|
||||
|
||||
# Request timeout in seconds (default: 30)
|
||||
# DISCORDBOT_COMFYUI_TIMEOUT=30
|
||||
|
||||
# =============================================================================
|
||||
# OPTIONAL - Database
|
||||
# =============================================================================
|
||||
|
||||
# Database URL (default: SQLite in bot/data/bot.db)
|
||||
# For PostgreSQL: postgresql+asyncpg://user:pass@localhost/dbname
|
||||
# DISCORDBOT_DATABASE_URL=sqlite+aiosqlite:///bot/data/bot.db
|
||||
|
||||
# =============================================================================
|
||||
# OPTIONAL - Defaults
|
||||
# =============================================================================
|
||||
|
||||
# Maximum pending jobs per user (default: 3)
|
||||
# DISCORDBOT_MAX_QUEUE_PER_USER=3
|
||||
|
||||
# Path to default workflow JSON file
|
||||
# DISCORDBOT_WORKFLOW_PATH=bot/data/default_workflow_api.json
|
||||
|
||||
# =============================================================================
|
||||
# OPTIONAL - Discord Application
|
||||
# =============================================================================
|
||||
|
||||
# Discord application ID (for slash command registration)
|
||||
# DISCORDBOT_APPLICATION_ID=your_app_id_here
|
||||
@@ -0,0 +1,115 @@
|
||||
name: Update Clone Count Badge
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: '0 0 * * *' # Runs daily at midnight
|
||||
workflow_dispatch: # Allows manual trigger
|
||||
|
||||
jobs:
|
||||
update-badge:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write # Needed to push to the branch
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Fetch Traffic Stats and Save JSON
|
||||
uses: actions/github-script@v6
|
||||
with:
|
||||
github-token: ${{ secrets.TRAFFIC_TOKEN }}
|
||||
script: |
|
||||
try {
|
||||
const { owner, repo } = context.repo;
|
||||
|
||||
// 1. Fetch Clones
|
||||
const clones = await github.rest.repos.getClones({
|
||||
owner,
|
||||
repo,
|
||||
});
|
||||
|
||||
// 2. Fetch Views (Visitors)
|
||||
const views = await github.rest.repos.getViews({
|
||||
owner,
|
||||
repo,
|
||||
});
|
||||
|
||||
// 3. Fetch Releases for Smart Download Count
|
||||
// We need to iterate through all releases to get the total count
|
||||
let smartDownloadCount = 0;
|
||||
let page = 1;
|
||||
let releases = [];
|
||||
|
||||
do {
|
||||
const response = await github.rest.repos.listReleases({
|
||||
owner,
|
||||
repo,
|
||||
per_page: 100,
|
||||
page: page
|
||||
});
|
||||
releases = response.data;
|
||||
|
||||
for (const release of releases) {
|
||||
let maxDownloads = 0;
|
||||
for (const asset of release.assets) {
|
||||
if (asset.download_count > maxDownloads) {
|
||||
maxDownloads = asset.download_count;
|
||||
}
|
||||
}
|
||||
smartDownloadCount += maxDownloads;
|
||||
}
|
||||
page++;
|
||||
} while (releases.length === 100);
|
||||
|
||||
const fs = require('fs');
|
||||
const data = {
|
||||
clones: {
|
||||
count: clones.data.count,
|
||||
uniques: clones.data.uniques,
|
||||
},
|
||||
views: {
|
||||
count: views.data.count,
|
||||
uniques: views.data.uniques,
|
||||
},
|
||||
downloads: {
|
||||
smart_count: smartDownloadCount
|
||||
},
|
||||
timestamp: new Date().toISOString()
|
||||
};
|
||||
|
||||
console.log('Stats fetched:', JSON.stringify(data));
|
||||
|
||||
// Write to a generic stats file
|
||||
fs.writeFileSync('traffic_stats.json', JSON.stringify(data, null, 2));
|
||||
// Maintain the old file for backward compatibility if needed, or just switch everything
|
||||
fs.writeFileSync('git_clones.json', JSON.stringify(data.clones, null, 2));
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error fetching stats:', error);
|
||||
process.exit(1);
|
||||
}
|
||||
|
||||
- name: Push to Badges Branch
|
||||
run: |
|
||||
git config --global user.name 'github-actions[bot]'
|
||||
git config --global user.email 'github-actions[bot]@users.noreply.github.com'
|
||||
|
||||
# Stash the files
|
||||
mv traffic_stats.json /tmp/traffic_stats.json
|
||||
mv git_clones.json /tmp/git_clones.json
|
||||
|
||||
# Fetch the badges branch or create it
|
||||
git fetch origin badges:badges || true
|
||||
git checkout badges || git checkout --orphan badges
|
||||
|
||||
# Clean the branch to ensure it only has the JSON
|
||||
git rm -rf .
|
||||
|
||||
# Restore the files
|
||||
mv /tmp/traffic_stats.json traffic_stats.json
|
||||
mv /tmp/git_clones.json git_clones.json
|
||||
|
||||
# Commit and push
|
||||
git add traffic_stats.json git_clones.json
|
||||
git diff --quiet && git diff --staged --quiet || (git commit -m "Update traffic statistics" && git push origin badges)
|
||||
+29
@@ -1,10 +1,39 @@
|
||||
# Sensitive files - NEVER commit these
|
||||
.env
|
||||
.env.*
|
||||
!.env.example
|
||||
config.yaml
|
||||
config.yml
|
||||
*.pem
|
||||
*.key
|
||||
secrets.*
|
||||
credentials.*
|
||||
|
||||
# Python
|
||||
__pycache__/
|
||||
*.pyc
|
||||
Errors.md
|
||||
Claude_Last_Convo.md
|
||||
ANALYSIS.md
|
||||
REPO_ANALYSIS.md
|
||||
PROJECT_BREAKDOWN.md
|
||||
REFACTOR_PRD.md
|
||||
PRD.md
|
||||
REPO_BREAKDOWN.md
|
||||
.jules/
|
||||
.Jules/
|
||||
.claude/
|
||||
.venv/
|
||||
venv/
|
||||
|
||||
# Standard Python/Test Artifacts
|
||||
.pytest_cache/
|
||||
htmlcov/
|
||||
.coverage
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
.mypy_cache/
|
||||
|
||||
# Likely test artifact
|
||||
MagicMock/
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
## 2026-01-14 - BytesIO Stream Position
|
||||
**Learning:** `img.save(bytes_io, ...)` writes to the buffer and leaves the cursor at the end. Subsequent reads return 0 bytes unless `bytes_io.seek(0)` is called.
|
||||
**Action:** Always verify stream position when working with in-memory buffers before passing them to IO-consuming functions.
|
||||
|
||||
## 2026-01-14 - Image Encoding Performance (Bolt Optimization)
|
||||
**Learning:**
|
||||
- **OpenCV** (`cv2.imencode`) is **~3x faster** than Pillow (`img.save`) for **PNG** encoding.
|
||||
- **Pillow** is **~30% faster** than OpenCV for **JPEG** encoding and avoids extra numpy conversion overhead.
|
||||
- **PyTorch tensor operations** (`(tensor * 255).to(uint8).numpy()`) are **~70% faster** than naive `numpy` conversion (`tensor.numpy() * 255`) by avoiding large float64 intermediate arrays.
|
||||
**Action:** Use PyTorch for tensor preprocessing. Use OpenCV for PNG, Pillow for JPEG/WebP.
|
||||
@@ -1,4 +0,0 @@
|
||||
## 2025-01-26 - Critical SSRF in Webhook Client
|
||||
**Vulnerability:** `send_to_discord_with_retry` accepted arbitrary URLs, allowing Server-Side Request Forgery (SSRF). A malicious user could probe internal services or cloud metadata services.
|
||||
**Learning:** The validation function `validate_webhook_url` existed but was not called in the main sending function. Also, `validate_webhook_url` had a fallback lenient check that could be bypassed.
|
||||
**Prevention:** Always enforce input validation at the point of use. Avoid "lenient" fallback checks for security-critical inputs like URLs.
|
||||
@@ -2,6 +2,43 @@
|
||||
|
||||
All notable changes to this project will be documented in this file.
|
||||
|
||||
## [2.0.0] - 2026-01-20
|
||||
|
||||
### Major Refactoring Release
|
||||
|
||||
Complete architectural refactoring to improve code organization, reduce duplication, and add bot features.
|
||||
|
||||
### Added
|
||||
- **Directory Structure**: New `nodes/`, `shared/`, `bot/` organization
|
||||
- **BaseDiscordNode**: Shared base class for image and video nodes (343 lines of reusable code)
|
||||
- **Shared Utilities**: 17 modular utility files in `shared/` package
|
||||
- `shared/discord/` - webhook client, message builder, CDN extractor
|
||||
- `shared/media/` - video encoder, format utils, image processing
|
||||
- `shared/workflow/` - sanitizer, prompt extractor, workflow builder
|
||||
- **Bot Features**:
|
||||
- WebSocket reconnection with exponential backoff
|
||||
- Error delivery to Discord users
|
||||
- `/template` commands (save, load, list, delete)
|
||||
- `/history` and `/rerun` commands
|
||||
- Config templates: `config.yaml.example`, `.env.example`
|
||||
|
||||
### Changed
|
||||
- **Code Reduction**: Total node code reduced by 620 lines (24%)
|
||||
- Image node: 986 → 836 lines
|
||||
- Video node: 1562 → 1092 lines
|
||||
- **Imports**: All imports now use `shared/` package instead of `discordsend_utils/`
|
||||
- **Video Encoding**: Extracted to `FFmpegEncoder` and `PILEncoder` classes
|
||||
|
||||
### Fixed
|
||||
- Bot startup bugs (BotConfig import, missing json import, permission imports)
|
||||
- SDXL workflow prompt extraction support
|
||||
- CDN URL redundant sends on 204 responses
|
||||
- Message builder metadata section formatting
|
||||
|
||||
### Removed
|
||||
- `discordsend_utils/` directory (replaced by `shared/`)
|
||||
- Obsolete documentation files
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### Changed
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
# ComfyUI-DiscordSend
|
||||
|
||||

|
||||

|
||||

|
||||

|
||||

|
||||
<br>
|
||||
@@ -24,12 +24,14 @@
|
||||
|
||||
---
|
||||
|
||||
## What's New in v1.1.0 (January 10, 2026)
|
||||
## What's New in v2.0.0 (January 20, 2026)
|
||||
|
||||
### 🚀 Core Updates
|
||||
- **Structured Logging**: Implemented comprehensive logging for better debugging and stability.
|
||||
- **Testing Suite**: Added initial test framework to ensure reliability.
|
||||
- **Enhanced Stability**: Various improvements to image and video handling logic.
|
||||
### 🚀 Major Refactoring Release
|
||||
- **New Architecture**: Reorganized into `nodes/`, `shared/`, and `bot/` packages for cleaner separation of concerns.
|
||||
- **BaseDiscordNode**: Shared base class eliminates 620 lines of duplicated code across image and video nodes.
|
||||
- **Shared Utilities**: 17 modular utility files covering Discord webhooks, media processing, and workflow handling.
|
||||
- **Discord Bot** *(optional)*: New standalone bot with slash commands, WebSocket reconnection, and job queuing.
|
||||
- **Performance**: Optimized image processing with direct Torch operations (~70% faster tensor processing).
|
||||
|
||||
📄 See [CHANGELOG.md](CHANGELOG.md) for the complete version history.
|
||||
|
||||
@@ -112,12 +114,33 @@
|
||||
cd /path/to/ComfyUI/custom_nodes
|
||||
git clone https://github.com/AEmotionStudio/ComfyUI-DiscordSend
|
||||
cd ComfyUI-DiscordSend
|
||||
pip install -r requirements.txt # Installs the minimal requirements (only the requests library)
|
||||
|
||||
# For nodes only (minimal - just the requests library):
|
||||
pip install -r requirements-nodes.txt
|
||||
|
||||
# For full bot support (Discord bot + all features):
|
||||
pip install -r requirements-bot.txt
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> - For video functionality, ffmpeg must be installed on your system. The node will automatically detect its presence.
|
||||
> - This extension has minimal dependencies, requiring only the 'requests' library which is included in the requirements.txt file.
|
||||
> - **Nodes only** require just the `requests` library (1 dependency).
|
||||
> - **Discord bot** requires additional dependencies (discord.py, aiohttp, sqlalchemy, etc.).
|
||||
|
||||
## 🏗️ Architecture Overview
|
||||
|
||||
This extension contains **two independent systems**. Most users only need the nodes.
|
||||
|
||||
| | ComfyUI Nodes | Discord Bot |
|
||||
|---|---|---|
|
||||
| **What it does** | Send images/videos to Discord from your workflow | Standalone bot that lets Discord users trigger ComfyUI workflows via slash commands |
|
||||
| **How to use** | Add `DiscordSendSaveImage` or `DiscordSendSaveVideo` to your workflow | Run `python -m bot` as a separate process alongside ComfyUI |
|
||||
| **Auth method** | Paste a **webhook URL** directly into the node | Requires a **bot token** from the Discord Developer Portal |
|
||||
| **Config files needed** | ❌ None — everything is configured in the node itself | ✅ `.env` or `config.yaml` (copy from `.env.example` / `config.yaml.example`) |
|
||||
| **Dependencies** | `requests` only | `discord.py`, `aiohttp`, `sqlalchemy`, etc. |
|
||||
|
||||
> [!NOTE]
|
||||
> The `.env.example` and `config.yaml.example` files in the repository root are **only for the optional Discord bot**. If you just want to send images/videos to Discord from your ComfyUI workflows, you do not need these files. Simply paste your Discord webhook URL into the node's `webhook_url` field.
|
||||
|
||||
## ⚙️ Settings
|
||||
|
||||
|
||||
+3
-3
@@ -10,9 +10,9 @@ current_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
if current_dir not in sys.path:
|
||||
sys.path.insert(0, current_dir)
|
||||
|
||||
# Import nodes
|
||||
from discord_image_node import DiscordSendSaveImage
|
||||
from discord_video_node import DiscordSendSaveVideo
|
||||
# Import nodes from nodes package using relative imports
|
||||
from .nodes.image_node import DiscordSendSaveImage
|
||||
from .nodes.video_node import DiscordSendSaveVideo
|
||||
|
||||
# Node class mappings for ComfyUI
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
+3
-3
@@ -8,9 +8,9 @@ from pathlib import Path
|
||||
project_root = Path(__file__).resolve().parent.parent
|
||||
sys.path.append(str(project_root))
|
||||
|
||||
from bot.config import Config
|
||||
from bot.config import BotConfig
|
||||
from bot.bot import ComfyUIBot
|
||||
from discordsend_utils.logging_config import setup_logging
|
||||
from shared.logging_config import setup_logging
|
||||
|
||||
def main():
|
||||
# Setup logging
|
||||
@@ -19,7 +19,7 @@ def main():
|
||||
|
||||
# Load configuration
|
||||
try:
|
||||
config = Config()
|
||||
config = BotConfig.load()
|
||||
except Exception as e:
|
||||
logger.critical(f"Failed to load configuration: {e}")
|
||||
return
|
||||
|
||||
+6
-7
@@ -3,9 +3,10 @@ from discord.ext import commands
|
||||
import logging
|
||||
import sys
|
||||
import asyncio
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
from .config import Config
|
||||
from .config import BotConfig
|
||||
from .database.repository import Repository
|
||||
from .comfyui.client import ComfyUIClient
|
||||
from .comfyui.websocket import ComfyUIWebSocket
|
||||
@@ -17,7 +18,7 @@ class ComfyUIBot(commands.Bot):
|
||||
Main Bot Class for ComfyUI Companion.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
def __init__(self, config: BotConfig):
|
||||
intents = discord.Intents.default()
|
||||
intents.message_content = True # Needed for some commands if not pure slash
|
||||
intents.members = True # Useful for permission checks
|
||||
@@ -32,9 +33,7 @@ class ComfyUIBot(commands.Bot):
|
||||
|
||||
# Database
|
||||
self.repository = Repository(config.database.url)
|
||||
|
||||
|
||||
import uuid
|
||||
|
||||
self.client_id = str(uuid.uuid4())
|
||||
|
||||
# ComfyUI Clients
|
||||
@@ -92,8 +91,8 @@ class ComfyUIBot(commands.Bot):
|
||||
extensions = [
|
||||
"bot.cogs.generate",
|
||||
"bot.cogs.queue",
|
||||
# "bot.cogs.templates",
|
||||
# "bot.cogs.history",
|
||||
"bot.cogs.templates",
|
||||
"bot.cogs.history",
|
||||
"bot.cogs.admin",
|
||||
]
|
||||
|
||||
|
||||
+2
-2
@@ -3,7 +3,7 @@ from discord import app_commands
|
||||
from discord.ext import commands
|
||||
import logging
|
||||
|
||||
from ..services.permissions import require_permission, Permissions, PermissionLevel
|
||||
from ..services.permissions import require_permission, Permissions
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -25,7 +25,7 @@ class AdminCog(commands.Cog):
|
||||
status_emoji = "✅" if comfy_status else "❌"
|
||||
|
||||
embed = discord.Embed(title="Bot Status", color=discord.Color.dark_grey())
|
||||
embed.add_field(name="ComfyUI Connection", value=f"{status_emoji} {self.bot.config.comfyui_url}", inline=False)
|
||||
embed.add_field(name="ComfyUI Connection", value=f"{status_emoji} {self.bot.config.comfyui.url}", inline=False)
|
||||
embed.add_field(name="Guilds", value=str(len(self.bot.guilds)), inline=True)
|
||||
embed.add_field(name="Latency", value=f"{round(self.bot.latency * 1000)}ms", inline=True)
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ from pathlib import Path
|
||||
|
||||
from ..embeds.builders import EmbedBuilder
|
||||
from ..services.permissions import require_permission, Permissions
|
||||
from ...discordsend_utils.workflow_builder import WorkflowBuilder
|
||||
from shared.workflow import WorkflowBuilder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
"""History cog for viewing past generations and rerunning them."""
|
||||
|
||||
import discord
|
||||
from discord import app_commands
|
||||
from discord.ext import commands
|
||||
import logging
|
||||
import json
|
||||
import random
|
||||
from typing import List
|
||||
|
||||
from ..services.permissions import require_permission, Permissions
|
||||
from ..database.models import JobStatus
|
||||
from ..embeds.builders import EmbedBuilder
|
||||
from shared.workflow import WorkflowBuilder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HistoryPaginator(discord.ui.View):
|
||||
"""Paginated view for job history."""
|
||||
|
||||
def __init__(self, jobs: List, per_page: int = 5):
|
||||
super().__init__(timeout=180)
|
||||
self.jobs = jobs
|
||||
self.per_page = per_page
|
||||
self.page = 0
|
||||
self.max_page = (len(jobs) - 1) // per_page if jobs else 0
|
||||
self._update_buttons()
|
||||
|
||||
def _update_buttons(self):
|
||||
self.prev_button.disabled = self.page <= 0
|
||||
self.next_button.disabled = self.page >= self.max_page
|
||||
|
||||
def get_embed(self) -> discord.Embed:
|
||||
embed = discord.Embed(title="Generation History", color=discord.Color.blue())
|
||||
|
||||
start = self.page * self.per_page
|
||||
end = start + self.per_page
|
||||
page_jobs = self.jobs[start:end]
|
||||
|
||||
if not page_jobs:
|
||||
embed.description = "No generation history found."
|
||||
return embed
|
||||
|
||||
lines = []
|
||||
for job in page_jobs:
|
||||
status_emoji = {
|
||||
JobStatus.COMPLETED.value: "✅",
|
||||
JobStatus.FAILED.value: "❌",
|
||||
JobStatus.CANCELLED.value: "🚫",
|
||||
JobStatus.PENDING.value: "⏳",
|
||||
JobStatus.RUNNING.value: "🔄",
|
||||
}.get(job.status, "❓")
|
||||
|
||||
prompt_preview = (job.positive_prompt or "No prompt")[:50]
|
||||
if len(job.positive_prompt or "") > 50:
|
||||
prompt_preview += "..."
|
||||
|
||||
timestamp = (
|
||||
job.created_at.strftime("%Y-%m-%d %H:%M") if job.created_at else "Unknown"
|
||||
)
|
||||
|
||||
lines.append(
|
||||
f"{status_emoji} **ID: {job.id}** | {timestamp}\n└ {prompt_preview}"
|
||||
)
|
||||
|
||||
embed.description = "\n\n".join(lines)
|
||||
embed.set_footer(
|
||||
text=f"Page {self.page + 1}/{self.max_page + 1} | Use /rerun <id> to regenerate"
|
||||
)
|
||||
|
||||
return embed
|
||||
|
||||
@discord.ui.button(label="◀ Previous", style=discord.ButtonStyle.secondary)
|
||||
async def prev_button(
|
||||
self, interaction: discord.Interaction, button: discord.ui.Button
|
||||
):
|
||||
self.page = max(0, self.page - 1)
|
||||
self._update_buttons()
|
||||
await interaction.response.edit_message(embed=self.get_embed(), view=self)
|
||||
|
||||
@discord.ui.button(label="Next ▶", style=discord.ButtonStyle.secondary)
|
||||
async def next_button(
|
||||
self, interaction: discord.Interaction, button: discord.ui.Button
|
||||
):
|
||||
self.page = min(self.max_page, self.page + 1)
|
||||
self._update_buttons()
|
||||
await interaction.response.edit_message(embed=self.get_embed(), view=self)
|
||||
|
||||
|
||||
class HistoryCog(commands.Cog):
|
||||
"""View generation history and rerun past jobs."""
|
||||
|
||||
def __init__(self, bot):
|
||||
self.bot = bot
|
||||
|
||||
@app_commands.command(name="history", description="View your generation history")
|
||||
@app_commands.describe(limit="Number of jobs to show (default: 20, max: 50)")
|
||||
async def history(self, interaction: discord.Interaction, limit: int = 20):
|
||||
"""Show paginated generation history."""
|
||||
await interaction.response.defer(ephemeral=True)
|
||||
|
||||
limit = min(max(1, limit), 50) # Clamp between 1 and 50
|
||||
|
||||
jobs = await self.bot.repository.list_user_jobs(
|
||||
user_discord_id=str(interaction.user.id), limit=limit
|
||||
)
|
||||
|
||||
if not jobs:
|
||||
await interaction.followup.send(
|
||||
"You have no generation history.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
view = HistoryPaginator(jobs)
|
||||
await interaction.followup.send(embed=view.get_embed(), view=view, ephemeral=True)
|
||||
|
||||
@app_commands.command(name="rerun", description="Rerun a previous generation")
|
||||
@app_commands.describe(job_id="The ID of the job to rerun")
|
||||
@require_permission(Permissions.USER.value)
|
||||
async def rerun(self, interaction: discord.Interaction, job_id: int):
|
||||
"""Rerun a previous job with the same parameters."""
|
||||
await interaction.response.defer()
|
||||
|
||||
# Get the original job
|
||||
original_job = await self.bot.repository.get_job_by_id(job_id)
|
||||
|
||||
if not original_job:
|
||||
await interaction.followup.send(
|
||||
f"Job ID {job_id} not found.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
# Check ownership
|
||||
if str(original_job.user.discord_id) != str(interaction.user.id):
|
||||
await interaction.followup.send(
|
||||
"You can only rerun your own jobs.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
# Check if job has required data
|
||||
if not original_job.workflow_json:
|
||||
await interaction.followup.send(
|
||||
"This job cannot be rerun (workflow data not saved).", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
workflow = json.loads(original_job.workflow_json)
|
||||
except json.JSONDecodeError:
|
||||
await interaction.followup.send(
|
||||
"Failed to parse original workflow.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
# Parse parameters
|
||||
parameters = {}
|
||||
if original_job.parameters:
|
||||
try:
|
||||
parameters = json.loads(original_job.parameters)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# Generate new seed for rerun
|
||||
new_seed = random.randint(1, 1000000000000000)
|
||||
parameters["seed"] = new_seed
|
||||
|
||||
# Update workflow with new seed
|
||||
builder = WorkflowBuilder(workflow)
|
||||
builder.set_seed(new_seed)
|
||||
final_workflow = builder.get_workflow()
|
||||
|
||||
# Determine delivery and context
|
||||
server_id = str(interaction.guild_id) if interaction.guild else None
|
||||
channel_id = str(interaction.channel_id)
|
||||
|
||||
try:
|
||||
# Create new job
|
||||
job = await self.bot.job_manager.create_job(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
workflow=final_workflow,
|
||||
positive_prompt=original_job.positive_prompt or "",
|
||||
negative_prompt=original_job.negative_prompt or "",
|
||||
parameters=parameters,
|
||||
server_discord_id=server_id,
|
||||
channel_id=channel_id,
|
||||
delivery_type=original_job.delivery_type or "channel",
|
||||
)
|
||||
|
||||
embed = EmbedBuilder.job_queued(job)
|
||||
embed.set_footer(text=f"Rerun of job #{job_id} | New job ID: {job.id}")
|
||||
|
||||
await interaction.followup.send(embed=embed)
|
||||
|
||||
# Store message ID
|
||||
original_message = await interaction.original_response()
|
||||
await self.bot.repository.update_job_message(
|
||||
job.prompt_id, str(original_message.id)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to rerun job: {e}")
|
||||
await interaction.followup.send(
|
||||
f"Failed to rerun job: {str(e)}", ephemeral=True
|
||||
)
|
||||
|
||||
|
||||
async def setup(bot):
|
||||
await bot.add_cog(HistoryCog(bot))
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Template management cog for saving and loading prompt presets."""
|
||||
|
||||
import discord
|
||||
from discord import app_commands
|
||||
from discord.ext import commands
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
from ..services.permissions import require_permission, Permissions
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TemplateCog(commands.Cog):
|
||||
"""Manage prompt templates."""
|
||||
|
||||
def __init__(self, bot):
|
||||
self.bot = bot
|
||||
|
||||
template_group = app_commands.Group(
|
||||
name="template", description="Manage prompt templates"
|
||||
)
|
||||
|
||||
@template_group.command(name="save", description="Save a prompt as a template")
|
||||
@app_commands.describe(
|
||||
name="Template name (unique per user/server)",
|
||||
prompt="The positive prompt to save",
|
||||
negative_prompt="Negative prompt (optional)",
|
||||
shared="Share with entire server (default: private)",
|
||||
)
|
||||
@require_permission(Permissions.GENERATOR.value)
|
||||
async def template_save(
|
||||
self,
|
||||
interaction: discord.Interaction,
|
||||
name: str,
|
||||
prompt: str,
|
||||
negative_prompt: str = "",
|
||||
shared: bool = False,
|
||||
):
|
||||
"""Save current prompt as a named template."""
|
||||
await interaction.response.defer(ephemeral=True)
|
||||
|
||||
# Validate name
|
||||
if len(name) > 100:
|
||||
await interaction.followup.send(
|
||||
"Template name must be 100 characters or less.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
# Ensure user exists
|
||||
await self.bot.repository.get_or_create_user(
|
||||
str(interaction.user.id), interaction.user.display_name
|
||||
)
|
||||
|
||||
server_id = str(interaction.guild_id) if interaction.guild and shared else None
|
||||
if server_id:
|
||||
await self.bot.repository.get_or_create_server(
|
||||
server_id, interaction.guild.name
|
||||
)
|
||||
|
||||
try:
|
||||
await self.bot.repository.create_template(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
name=name,
|
||||
positive_prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
server_discord_id=server_id,
|
||||
)
|
||||
|
||||
scope = "server" if shared else "private"
|
||||
await interaction.followup.send(
|
||||
f"Saved template **{name}** ({scope}).", ephemeral=True
|
||||
)
|
||||
except Exception as e:
|
||||
if "UNIQUE constraint" in str(e):
|
||||
await interaction.followup.send(
|
||||
f"A template named **{name}** already exists. "
|
||||
"Delete it first or use a different name.",
|
||||
ephemeral=True,
|
||||
)
|
||||
else:
|
||||
logger.error(f"Failed to save template: {e}")
|
||||
await interaction.followup.send(
|
||||
"Failed to save template.", ephemeral=True
|
||||
)
|
||||
|
||||
@template_group.command(name="load", description="Load a saved template")
|
||||
@app_commands.describe(name="Template name to load")
|
||||
async def template_load(self, interaction: discord.Interaction, name: str):
|
||||
"""Load a template and show its contents."""
|
||||
await interaction.response.defer(ephemeral=True)
|
||||
|
||||
server_id = str(interaction.guild_id) if interaction.guild else None
|
||||
|
||||
# Try user's private template first
|
||||
template = await self.bot.repository.get_template(
|
||||
user_discord_id=str(interaction.user.id), name=name
|
||||
)
|
||||
|
||||
# Try shared server template if not found
|
||||
if not template and server_id:
|
||||
templates = await self.bot.repository.list_templates(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
server_discord_id=server_id,
|
||||
include_shared=True,
|
||||
)
|
||||
template = next((t for t in templates if t.name == name), None)
|
||||
|
||||
if not template:
|
||||
await interaction.followup.send(
|
||||
f"Template **{name}** not found.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
embed = discord.Embed(title=f"Template: {template.name}", color=discord.Color.blue())
|
||||
embed.add_field(
|
||||
name="Prompt", value=template.positive_prompt[:1024], inline=False
|
||||
)
|
||||
if template.negative_prompt:
|
||||
embed.add_field(
|
||||
name="Negative Prompt",
|
||||
value=template.negative_prompt[:1024],
|
||||
inline=False,
|
||||
)
|
||||
|
||||
scope = "Shared" if template.server_id else "Private"
|
||||
embed.set_footer(text=f"{scope} template | Use /generate with this prompt")
|
||||
|
||||
await interaction.followup.send(embed=embed, ephemeral=True)
|
||||
|
||||
@template_group.command(name="list", description="List your saved templates")
|
||||
async def template_list(self, interaction: discord.Interaction):
|
||||
"""List all available templates."""
|
||||
await interaction.response.defer(ephemeral=True)
|
||||
|
||||
server_id = str(interaction.guild_id) if interaction.guild else None
|
||||
|
||||
templates = await self.bot.repository.list_templates(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
server_discord_id=server_id,
|
||||
include_shared=True,
|
||||
)
|
||||
|
||||
if not templates:
|
||||
await interaction.followup.send(
|
||||
"You have no saved templates.", ephemeral=True
|
||||
)
|
||||
return
|
||||
|
||||
embed = discord.Embed(title="Your Templates", color=discord.Color.blue())
|
||||
|
||||
private_templates = [t for t in templates if t.server_id is None]
|
||||
shared_templates = [t for t in templates if t.server_id is not None]
|
||||
|
||||
if private_templates:
|
||||
names = "\n".join([f"• {t.name}" for t in private_templates[:10]])
|
||||
if len(private_templates) > 10:
|
||||
names += f"\n...and {len(private_templates) - 10} more"
|
||||
embed.add_field(name="Private Templates", value=names, inline=False)
|
||||
|
||||
if shared_templates:
|
||||
names = "\n".join([f"• {t.name}" for t in shared_templates[:10]])
|
||||
if len(shared_templates) > 10:
|
||||
names += f"\n...and {len(shared_templates) - 10} more"
|
||||
embed.add_field(name="Server Templates", value=names, inline=False)
|
||||
|
||||
await interaction.followup.send(embed=embed, ephemeral=True)
|
||||
|
||||
@template_group.command(name="delete", description="Delete a saved template")
|
||||
@app_commands.describe(name="Template name to delete")
|
||||
@require_permission(Permissions.GENERATOR.value)
|
||||
async def template_delete(self, interaction: discord.Interaction, name: str):
|
||||
"""Delete a template."""
|
||||
await interaction.response.defer(ephemeral=True)
|
||||
|
||||
# Try deleting private template
|
||||
deleted = await self.bot.repository.delete_template(
|
||||
user_discord_id=str(interaction.user.id), name=name
|
||||
)
|
||||
|
||||
# Try deleting shared template if private not found
|
||||
if not deleted and interaction.guild:
|
||||
deleted = await self.bot.repository.delete_template(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
name=name,
|
||||
server_discord_id=str(interaction.guild_id),
|
||||
)
|
||||
|
||||
if deleted:
|
||||
await interaction.followup.send(
|
||||
f"Deleted template **{name}**.", ephemeral=True
|
||||
)
|
||||
else:
|
||||
await interaction.followup.send(
|
||||
f"Template **{name}** not found or you don't have permission to delete it.",
|
||||
ephemeral=True,
|
||||
)
|
||||
|
||||
@template_load.autocomplete("name")
|
||||
@template_delete.autocomplete("name")
|
||||
async def template_name_autocomplete(
|
||||
self, interaction: discord.Interaction, current: str
|
||||
) -> List[app_commands.Choice[str]]:
|
||||
"""Autocomplete for template names."""
|
||||
server_id = str(interaction.guild_id) if interaction.guild else None
|
||||
|
||||
templates = await self.bot.repository.list_templates(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
server_discord_id=server_id,
|
||||
include_shared=True,
|
||||
)
|
||||
|
||||
# Filter by current input
|
||||
filtered = [t for t in templates if current.lower() in t.name.lower()]
|
||||
|
||||
return [
|
||||
app_commands.Choice(name=t.name, value=t.name)
|
||||
for t in filtered[:25] # Discord limit
|
||||
]
|
||||
|
||||
|
||||
async def setup(bot):
|
||||
await bot.add_cog(TemplateCog(bot))
|
||||
+128
-20
@@ -2,10 +2,17 @@ import aiohttp
|
||||
import logging
|
||||
import json
|
||||
import asyncio
|
||||
import random
|
||||
from typing import Callable, Coroutine, Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Reconnection constants
|
||||
INITIAL_BACKOFF = 1.0 # Initial delay in seconds
|
||||
MAX_BACKOFF = 60.0 # Maximum delay cap
|
||||
BACKOFF_MULTIPLIER = 2 # Exponential multiplier
|
||||
JITTER_FACTOR = 0.1 # +/- 10% randomization
|
||||
|
||||
class ComfyUIWebSocket:
|
||||
"""WebSocket Client for real-time ComfyUI events."""
|
||||
|
||||
@@ -27,25 +34,64 @@ class ComfyUIWebSocket:
|
||||
self._running = False
|
||||
self._listen_task: Optional[asyncio.Task] = None
|
||||
|
||||
async def connect(self):
|
||||
"""Connect to the WebSocket."""
|
||||
if self.session is None or self.session.closed:
|
||||
self.session = aiohttp.ClientSession()
|
||||
# Reconnection state
|
||||
self._reconnect_attempts = 0
|
||||
self._should_reconnect = True
|
||||
self._reconnect_task: Optional[asyncio.Task] = None
|
||||
|
||||
try:
|
||||
self.ws = await self.session.ws_connect(self.ws_url)
|
||||
self._running = True
|
||||
self._listen_task = asyncio.create_task(self._listen())
|
||||
logger.info(f"Connected to ComfyUI WebSocket at {self.ws_url}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to WebSocket: {e}")
|
||||
if self.session and not self.session.closed:
|
||||
await self.session.close()
|
||||
raise
|
||||
# Lock to prevent race conditions between connect() and disconnect()
|
||||
self._state_lock = asyncio.Lock()
|
||||
|
||||
async def connect(self) -> bool:
|
||||
"""Connect to the WebSocket.
|
||||
|
||||
Returns:
|
||||
True if connection was established successfully, False if aborted.
|
||||
"""
|
||||
async with self._state_lock:
|
||||
# Check if disconnect was called - don't proceed if so
|
||||
if not self._should_reconnect:
|
||||
logger.info("Connect aborted: disconnect was requested")
|
||||
return False
|
||||
|
||||
if self.session is None or self.session.closed:
|
||||
self.session = aiohttp.ClientSession()
|
||||
|
||||
try:
|
||||
self.ws = await self.session.ws_connect(self.ws_url)
|
||||
# Re-check after await in case disconnect() was called during connection
|
||||
if not self._should_reconnect:
|
||||
logger.info("Connect aborted after ws_connect: disconnect was requested")
|
||||
if self.ws and not self.ws.closed:
|
||||
await self.ws.close()
|
||||
if self.session and not self.session.closed:
|
||||
await self.session.close()
|
||||
return False
|
||||
self._running = True
|
||||
self._reconnect_attempts = 0 # Reset on successful connection
|
||||
self._listen_task = asyncio.create_task(self._listen())
|
||||
logger.info(f"Connected to ComfyUI WebSocket at {self.ws_url}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to WebSocket: {e}")
|
||||
if self.session and not self.session.closed:
|
||||
await self.session.close()
|
||||
raise
|
||||
|
||||
async def disconnect(self):
|
||||
"""Disconnect from WebSocket."""
|
||||
self._running = False
|
||||
async with self._state_lock:
|
||||
self._should_reconnect = False # Prevent reconnection loop
|
||||
self._running = False
|
||||
|
||||
# Cancel reconnection task if running (outside lock to avoid deadlock)
|
||||
if self._reconnect_task and not self._reconnect_task.done():
|
||||
self._reconnect_task.cancel()
|
||||
try:
|
||||
await self._reconnect_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
if self.ws:
|
||||
await self.ws.close()
|
||||
if self.session:
|
||||
@@ -89,9 +135,71 @@ class ComfyUIWebSocket:
|
||||
logger.error("WebSocket connection closed with error")
|
||||
break
|
||||
except Exception as e:
|
||||
if self._running:
|
||||
logger.error(f"WebSocket listener error: {e}")
|
||||
# Verify reconnection logic would handle this or let job manager handle it
|
||||
if self._running:
|
||||
logger.error(f"WebSocket listener error: {e}")
|
||||
finally:
|
||||
if self._running:
|
||||
logger.info("WebSocket listener stopped unexpectedly.")
|
||||
if self._running and self._should_reconnect:
|
||||
logger.warning("WebSocket connection lost. Starting reconnection...")
|
||||
self._reconnect_task = asyncio.create_task(self._handle_disconnect())
|
||||
|
||||
def _calculate_backoff(self) -> float:
|
||||
"""Calculate backoff delay with exponential growth and jitter."""
|
||||
delay = INITIAL_BACKOFF * (BACKOFF_MULTIPLIER ** self._reconnect_attempts)
|
||||
delay = min(delay, MAX_BACKOFF)
|
||||
# Add jitter: +/- JITTER_FACTOR
|
||||
jitter = delay * JITTER_FACTOR * (2 * random.random() - 1)
|
||||
return delay + jitter
|
||||
|
||||
async def _handle_disconnect(self) -> None:
|
||||
"""Handle unexpected disconnection by attempting to reconnect."""
|
||||
self._running = False
|
||||
|
||||
# Close existing connections
|
||||
if self.ws and not self.ws.closed:
|
||||
try:
|
||||
await self.ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
if self.session and not self.session.closed:
|
||||
try:
|
||||
await self.session.close()
|
||||
except Exception:
|
||||
pass
|
||||
self.session = None
|
||||
self.ws = None
|
||||
|
||||
await self._reconnect_loop()
|
||||
|
||||
async def _reconnect_loop(self) -> None:
|
||||
"""Background task that handles reconnection with exponential backoff."""
|
||||
while self._should_reconnect:
|
||||
self._reconnect_attempts += 1
|
||||
backoff = self._calculate_backoff()
|
||||
|
||||
logger.info(
|
||||
f"Reconnection attempt {self._reconnect_attempts} "
|
||||
f"in {backoff:.1f}s..."
|
||||
)
|
||||
await asyncio.sleep(backoff)
|
||||
|
||||
if not self._should_reconnect:
|
||||
logger.info("Reconnection cancelled.")
|
||||
return
|
||||
|
||||
try:
|
||||
attempts_made = self._reconnect_attempts
|
||||
connected = await self.connect()
|
||||
if connected:
|
||||
logger.info(
|
||||
f"Successfully reconnected after "
|
||||
f"{attempts_made} attempt(s)."
|
||||
)
|
||||
return
|
||||
else:
|
||||
# connect() returned False - disconnect was called
|
||||
logger.info("Reconnection aborted: disconnect was requested")
|
||||
return
|
||||
except Exception as e:
|
||||
logger.warning(f"Reconnection attempt failed: {e}")
|
||||
|
||||
logger.info("Reconnection loop ended (should_reconnect=False).")
|
||||
|
||||
+68
-17
@@ -1,12 +1,15 @@
|
||||
import discord
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
from typing import List, Dict, Any, Union
|
||||
from typing import Any, Optional, Union
|
||||
|
||||
from ..comfyui.client import ComfyUIClient
|
||||
from ..embeds.builders import EmbedBuilder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DeliveryService:
|
||||
"""Handles delivery of results to Discord."""
|
||||
|
||||
@@ -14,7 +17,69 @@ class DeliveryService:
|
||||
self.bot = bot
|
||||
self.client = comfy_client
|
||||
|
||||
async def deliver_job(self, job: Any): # job: Job model
|
||||
async def _get_destination(
|
||||
self, job: Any
|
||||
) -> Optional[Union[discord.User, discord.TextChannel]]:
|
||||
"""Get the destination channel or user for a job."""
|
||||
if job.delivery_type == "dm":
|
||||
try:
|
||||
user = self.bot.get_user(int(job.user.discord_id))
|
||||
if not user:
|
||||
user = await self.bot.fetch_user(int(job.user.discord_id))
|
||||
return user
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch user for DM: {e}")
|
||||
return None
|
||||
|
||||
if job.channel_id:
|
||||
destination = self.bot.get_channel(int(job.channel_id))
|
||||
if destination:
|
||||
return destination
|
||||
# Fallback to DM if channel not found
|
||||
try:
|
||||
user = self.bot.get_user(int(job.user.discord_id))
|
||||
if not user:
|
||||
user = await self.bot.fetch_user(int(job.user.discord_id))
|
||||
return user
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
async def deliver_error(self, job: Any, error_message: str) -> bool:
|
||||
"""
|
||||
Deliver error notification for a failed job.
|
||||
|
||||
Args:
|
||||
job: The failed Job model instance
|
||||
error_message: The error message to display
|
||||
|
||||
Returns:
|
||||
True if delivery succeeded, False otherwise
|
||||
"""
|
||||
# Truncate error message if too long (Discord embed field limit)
|
||||
if len(error_message) > 1000:
|
||||
error_message = error_message[:997] + "..."
|
||||
|
||||
embed = EmbedBuilder.job_failed(job, error_message)
|
||||
|
||||
destination = await self._get_destination(job)
|
||||
|
||||
if destination:
|
||||
try:
|
||||
await destination.send(embed=embed)
|
||||
logger.info(f"Delivered error for job {job.id} to {destination}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to deliver error: {e}")
|
||||
return False
|
||||
else:
|
||||
logger.error(
|
||||
f"Could not determine destination for error delivery (job {job.id})"
|
||||
)
|
||||
return False
|
||||
|
||||
async def deliver_job(self, job: Any): # job: Job model
|
||||
"""Deliver results for a completed job."""
|
||||
if not job.output_images:
|
||||
logger.warning(f"Job {job.id} completed but has no images.")
|
||||
@@ -42,21 +107,7 @@ class DeliveryService:
|
||||
logger.warning("No files to upload.")
|
||||
return
|
||||
|
||||
# Determine destination
|
||||
destination = None
|
||||
|
||||
if job.delivery_type == "dm":
|
||||
user = self.bot.get_user(int(job.user.discord_id)) or await self.bot.fetch_user(int(job.user.discord_id))
|
||||
destination = user
|
||||
elif job.channel_id:
|
||||
destination = self.bot.get_channel(int(job.channel_id))
|
||||
if not destination:
|
||||
# Fallback to DM if channel not found?
|
||||
try:
|
||||
user = self.bot.get_user(int(job.user.discord_id)) or await self.bot.fetch_user(int(job.user.discord_id))
|
||||
destination = user
|
||||
except:
|
||||
pass
|
||||
destination = await self._get_destination(job)
|
||||
|
||||
if destination:
|
||||
content = f"Generation complete for <@{job.user.discord_id}>!\n**Prompt:** {job.positive_prompt}"
|
||||
|
||||
@@ -196,11 +196,13 @@ class JobManager:
|
||||
prompt_id = msg.get("prompt_id")
|
||||
exception_type = msg.get("exception_type", "Unknown Error")
|
||||
exception_message = msg.get("exception_message", "")
|
||||
|
||||
|
||||
if prompt_id:
|
||||
error_msg = f"{exception_type}: {exception_message}"
|
||||
job = await self.repo.update_job_status(prompt_id, JobStatus.FAILED.value, error_message=error_msg)
|
||||
# Notify user of failure via delivery service?
|
||||
# Ideally yes, but DeliveryService currently only sends images.
|
||||
# We might want to expand DeliveryService to handle errors too, or reuse the channel_id to post the failure embed.
|
||||
pass
|
||||
job = await self.repo.update_job_status(
|
||||
prompt_id, JobStatus.FAILED.value, error_message=error_msg
|
||||
)
|
||||
|
||||
# Deliver error notification to user
|
||||
if job:
|
||||
await self.delivery.deliver_error(job, error_msg)
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
# ComfyUI-DiscordSend Bot Configuration
|
||||
# Copy this file to bot/config.yaml and fill in your values
|
||||
# Environment variables will override these settings (see .env.example)
|
||||
|
||||
discord:
|
||||
# Your Discord bot token (required)
|
||||
# Get one at: https://discord.com/developers/applications
|
||||
token: "YOUR_BOT_TOKEN_HERE"
|
||||
|
||||
# Your Discord application ID (optional)
|
||||
# Used for slash command registration
|
||||
application_id: "YOUR_APP_ID_HERE"
|
||||
|
||||
comfyui:
|
||||
# ComfyUI server URL (default: http://127.0.0.1:8188)
|
||||
url: "http://127.0.0.1:8188"
|
||||
|
||||
# WebSocket URL for real-time updates (auto-derived from url if not set)
|
||||
# ws_url: "ws://127.0.0.1:8188/ws"
|
||||
|
||||
# Request timeout in seconds
|
||||
timeout: 30
|
||||
|
||||
defaults:
|
||||
# Maximum pending jobs per user per server
|
||||
max_queue_per_user: 3
|
||||
|
||||
# How often to update progress embeds (seconds)
|
||||
progress_update_interval: 2.0
|
||||
|
||||
# Default generation parameters
|
||||
default_steps: 20
|
||||
default_cfg: 7.0
|
||||
default_width: 512
|
||||
default_height: 512
|
||||
|
||||
# Path to default workflow (relative to bot/data/)
|
||||
# workflow_path: "default_workflow_api.json"
|
||||
|
||||
database:
|
||||
# Database URL (default: SQLite in bot/data/bot.db)
|
||||
# For PostgreSQL: postgresql+asyncpg://user:pass@localhost/dbname
|
||||
# url: "sqlite+aiosqlite:///bot/data/bot.db"
|
||||
|
||||
security:
|
||||
# Restrict bot to specific guild IDs (empty = all guilds allowed)
|
||||
allowed_guilds: []
|
||||
# Example: allowed_guilds: [123456789, 987654321]
|
||||
@@ -1,993 +0,0 @@
|
||||
"""ComfyUI node for sending images to Discord and saving them locally."""
|
||||
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import torch
|
||||
import folder_paths
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
from comfy.cli_args import args
|
||||
import re
|
||||
import cv2
|
||||
import requests
|
||||
from io import BytesIO
|
||||
from uuid import uuid4
|
||||
from typing import Any, Union, List, Optional
|
||||
|
||||
# Import shared utilities
|
||||
from discordsend_utils import (
|
||||
sanitize_json_for_export,
|
||||
update_github_cdn_urls,
|
||||
extract_prompts_from_workflow,
|
||||
send_to_discord_with_retry
|
||||
)
|
||||
|
||||
|
||||
# Helper function to convert tensor to OpenCV format
|
||||
def tensor_to_cv(tensor: torch.Tensor) -> np.ndarray:
|
||||
"""Convert a PyTorch tensor to an OpenCV-compatible numpy array."""
|
||||
# Optimization: Use torch operations for scaling/clipping/casting to avoid large float64 intermediate arrays on CPU
|
||||
return (tensor.squeeze() * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
|
||||
|
||||
|
||||
class DiscordSendSaveImage:
|
||||
"""
|
||||
A ComfyUI node that can send images to Discord and save them with advanced options.
|
||||
Images can be sent to Discord via webhook integration, while providing flexible
|
||||
saving options with customizable naming conventions and format options.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.type = "output"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
self.output_dir = None # Will be set during saving to store the actual path used
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE", {"tooltip": "The images to save and/or send to Discord."}),
|
||||
"filename_prefix": ("STRING", {"default": "ComfyUI-Image", "tooltip": "The prefix for the saved files."}),
|
||||
"overwrite_last": ("BOOLEAN", {"default": False, "tooltip": "If enabled, will overwrite the last image instead of creating incrementing filenames."})
|
||||
},
|
||||
"optional": {
|
||||
"file_format": (["png", "jpeg", "webp"], {
|
||||
"default": "png",
|
||||
"tooltip": "The format to save images in. PNG is lossless but larger. JPEG and WebP are smaller but lossy."
|
||||
}),
|
||||
"quality": ("INT", {
|
||||
"default": 95,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"step": 1,
|
||||
"tooltip": "Quality (1-100) for JPEG/WebP. Ignored for PNG. Higher values = better quality but larger file size."
|
||||
}),
|
||||
"lossless": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "If enabled, will use lossless compression for supported formats (PNG and WebP). JPEG will use maximum quality."
|
||||
}),
|
||||
"save_output": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Whether to save images to disk. When disabled, images will only be previewed in the UI."
|
||||
}),
|
||||
"show_preview": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Whether to show image previews in the UI. Disable to reduce UI clutter for large batches."
|
||||
}),
|
||||
"add_date": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Add the current date (YYYY-MM-DD) to filenames."
|
||||
}),
|
||||
"add_time": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Add the current time (HH-MM-SS) to filenames."
|
||||
}),
|
||||
"add_dimensions": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Add width and height dimensions to the filename (WxH format)."
|
||||
}),
|
||||
"resize_to_power_of_2": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Resize images to nearest power of 2 dimensions (useful for textures in game engines)."
|
||||
}),
|
||||
"resize_method": (["nearest-exact", "bilinear", "bicubic", "lanczos", "box"], {
|
||||
"default": "lanczos",
|
||||
"tooltip": "The method to use when resizing images. Lanczos generally provides the best quality but may be slower."
|
||||
}),
|
||||
"send_to_discord": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to send the images to Discord via webhook."
|
||||
}),
|
||||
"webhook_url": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "Discord webhook URL (from Server Settings > Integrations > Webhooks). Leave empty to disable Discord integration."
|
||||
}),
|
||||
"discord_message": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional text to display with the image. Supports Discord Markdown (bold, italic, etc.)."
|
||||
}),
|
||||
"include_prompts_in_message": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to include the positive and negative prompts in the Discord message."
|
||||
}),
|
||||
"include_format_in_message": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to include the image format in the Discord message."
|
||||
}),
|
||||
"group_batched_images": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Group all images from a batch into a single Discord message with a gallery, rather than sending each one separately. Maximum is 9 images."
|
||||
}),
|
||||
"send_workflow_json": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to send the workflow JSON alongside the image to Discord, allowing dragging the JSON into ComfyUI to restore the workflow."
|
||||
}),
|
||||
"save_cdn_urls": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to save the Discord CDN URLs of the uploaded images as a text file and attach it to the Discord message."
|
||||
}),
|
||||
"github_cdn_update": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to update a GitHub repository with the Discord CDN URLs."
|
||||
}),
|
||||
"github_repo": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "GitHub repository to update with CDN URLs (format: username/repo)."
|
||||
}),
|
||||
"github_token": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "GitHub personal access token with repo permissions."
|
||||
}),
|
||||
"github_file_path": ("STRING", {
|
||||
"default": "cdn_urls.md",
|
||||
"multiline": False,
|
||||
"tooltip": "Path to the file within the GitHub repository to update with CDN URLs."
|
||||
}),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO"
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("image_path",)
|
||||
FUNCTION = "save_images"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "image/output"
|
||||
DESCRIPTION = "Saves images with advanced options and can send them to Discord via webhook integration. Returns the path to the first saved image."
|
||||
|
||||
@classmethod
|
||||
def CONTEXT_MENUS(s):
|
||||
return {
|
||||
"Show Preview": lambda self, **kwargs: {"show_preview": True},
|
||||
"Hide Preview": lambda self, **kwargs: {"show_preview": False},
|
||||
}
|
||||
|
||||
def save_images(self, images, filename_prefix="ComfyUI-Image", overwrite_last=False,
|
||||
file_format="png", quality=95, lossless=True, add_date=False, add_time=False,
|
||||
add_dimensions=False, resize_to_power_of_2=False, save_output=True,
|
||||
resize_method="lanczos", show_preview=True, send_to_discord=False, webhook_url="", discord_message="",
|
||||
include_prompts_in_message=False, include_format_in_message=False, send_workflow_json=False,
|
||||
group_batched_images=True, save_cdn_urls=False, github_cdn_update=False, github_repo="",
|
||||
github_token="", github_file_path="cdn_urls.md", prompt=None, extra_pnginfo=None):
|
||||
"""
|
||||
Save images for and optionally send to Discord.
|
||||
|
||||
Parameters:
|
||||
images: The images to save/send
|
||||
filename_prefix: The prefix for the filename
|
||||
overwrite_last: Whether to overwrite the last image instead of incrementing
|
||||
file_format: Image format to save as (png, jpeg, webp)
|
||||
quality: Quality setting for lossy formats (1-100)
|
||||
lossless: Whether to use lossless compression for supported formats (PNG and WebP)
|
||||
add_date: Whether to add the date to the filename
|
||||
add_time: Whether to add the time to the filename
|
||||
add_dimensions: Whether to add the image dimensions to the filename
|
||||
resize_to_power_of_2: Whether to resize to power-of-2 dimensions for texture optimization
|
||||
resize_method: Method to use for resizing
|
||||
save_output: Whether to save to disk or just preview
|
||||
send_to_discord: Whether to send the images to Discord
|
||||
webhook_url: Discord webhook URL
|
||||
discord_message: Message to send with the images
|
||||
include_prompts_in_message: Whether to include prompts in Discord message
|
||||
include_format_in_message: Whether to include the image format in Discord messages
|
||||
send_workflow_json: Whether to send the workflow JSON to Discord
|
||||
group_batched_images: Whether to group all images from a batch into a single Discord message
|
||||
save_cdn_urls: Whether to save Discord CDN URLs as a text file and attach it to the Discord message
|
||||
github_cdn_update: Whether to update a GitHub repository with the Discord CDN URLs
|
||||
github_repo: GitHub repository (format: username/repo)
|
||||
github_token: GitHub personal access token
|
||||
github_file_path: Path to the file within the GitHub repository to update
|
||||
prompt: The generation prompt data
|
||||
extra_pnginfo: Extra PNG info for metadata
|
||||
|
||||
Returns:
|
||||
UI information for ComfyUI and the path to the first saved image as a string.
|
||||
If no images were saved, an empty string is returned for the path.
|
||||
"""
|
||||
results = []
|
||||
output_files = []
|
||||
discord_sent_files = []
|
||||
discord_send_success = True
|
||||
|
||||
# For batch grouping
|
||||
batch_discord_files = []
|
||||
batch_discord_data = {}
|
||||
batch_workflow_json = None
|
||||
|
||||
# For tracking Discord CDN URLs
|
||||
discord_cdn_urls = []
|
||||
batch_cdn_urls = []
|
||||
|
||||
# Sanitize the workflow and extra_pnginfo data to remove webhook URLs
|
||||
# This protects user security when sharing images
|
||||
# (but keep a copy of the original data for prompt extraction)
|
||||
original_prompt = prompt
|
||||
original_extra_pnginfo = extra_pnginfo
|
||||
|
||||
# Ensure webhook URL is sanitized from workflow data for all file formats
|
||||
if prompt is not None:
|
||||
prompt = sanitize_json_for_export(prompt)
|
||||
|
||||
if extra_pnginfo is not None:
|
||||
extra_pnginfo = sanitize_json_for_export(extra_pnginfo)
|
||||
|
||||
# Double-check webhook URL removal for Discord-specific data
|
||||
if send_to_discord:
|
||||
# Verify webhook is sanitized from workflow JSON data
|
||||
if send_workflow_json and extra_pnginfo is not None and "workflow" in extra_pnginfo:
|
||||
extra_pnginfo["workflow"] = sanitize_json_for_export(extra_pnginfo["workflow"])
|
||||
|
||||
# Add date and/or time if enabled
|
||||
date_time_parts = []
|
||||
|
||||
# Prepare info for Discord message
|
||||
image_info = {}
|
||||
|
||||
if add_date:
|
||||
# Get ONLY the date in YYYY-MM-DD format
|
||||
current_date = time.strftime("%Y-%m-%d")
|
||||
date_time_parts.append(current_date)
|
||||
print(f"Adding date to filename: {current_date}")
|
||||
image_info["date"] = current_date
|
||||
|
||||
if add_time:
|
||||
# Get ONLY the time in HH-MM-SS format
|
||||
current_time = time.strftime("%H-%M-%S")
|
||||
date_time_parts.append(current_time)
|
||||
print(f"Adding time to filename: {current_time}")
|
||||
image_info["time"] = current_time
|
||||
|
||||
# Add date/time components to filename prefix if any were enabled
|
||||
if date_time_parts:
|
||||
date_time_suffix = "_" + "_".join(date_time_parts)
|
||||
filename_prefix += date_time_suffix
|
||||
print(f"Final timestamp suffix: {date_time_suffix}")
|
||||
|
||||
# Add prefix append
|
||||
filename_prefix += self.prefix_append
|
||||
|
||||
# Get ComfyUI output directory for safe path handling
|
||||
comfy_output_dir = folder_paths.get_output_directory()
|
||||
|
||||
# Choose destination directory based on save_output flag
|
||||
if save_output:
|
||||
# Create a output subfolder in the ComfyUI output directory
|
||||
dest_folder = os.path.join(comfy_output_dir, "discord_output")
|
||||
os.makedirs(dest_folder, exist_ok=True)
|
||||
else:
|
||||
# Use ComfyUI's temporary directory for preview-only files
|
||||
dest_folder = folder_paths.get_temp_directory()
|
||||
os.makedirs(dest_folder, exist_ok=True)
|
||||
print(f"Using temporary directory for preview: {dest_folder}")
|
||||
|
||||
# Setup paths using ComfyUI's path validation
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix, dest_folder, images[0].shape[1], images[0].shape[0])
|
||||
|
||||
# For overwrite functionality, we'll just always use the same counter instead of bypassing validation
|
||||
if overwrite_last == "enable":
|
||||
counter = 1 # Always use the same counter value for overwriting
|
||||
else:
|
||||
# When not overwriting, we need to find the highest existing counter and start from there
|
||||
# This ensures we're always creating new files
|
||||
try:
|
||||
# Get all existing files with this prefix
|
||||
base_filename = os.path.basename(filename).replace("%batch_num%", "")
|
||||
existing_files = [f for f in os.listdir(full_output_folder)
|
||||
if os.path.basename(f).startswith(base_filename)]
|
||||
|
||||
if existing_files:
|
||||
# Extract counters from filenames
|
||||
existing_counters = []
|
||||
for f in existing_files:
|
||||
# Extract counter pattern (5 digits) from filename
|
||||
counter_match = re.search(r'_(\d{5})\.', f)
|
||||
if counter_match:
|
||||
existing_counters.append(int(counter_match.group(1)))
|
||||
|
||||
# Also try alternative pattern where the counter is followed by extension
|
||||
counter_match = re.search(r'_(\d{5})_\.', f)
|
||||
if counter_match:
|
||||
existing_counters.append(int(counter_match.group(1)))
|
||||
|
||||
# Set counter to one more than the highest existing counter
|
||||
if existing_counters:
|
||||
counter = max(existing_counters) + 1
|
||||
except Exception as e:
|
||||
print(f"Error determining next file counter: {e}")
|
||||
# Default to ComfyUI's counter if we can't determine the next one
|
||||
|
||||
print(f"Using counter: {counter} for {'overwriting' if overwrite_last == 'enable' else 'new files'}")
|
||||
print(f"Output prefix: {filename_prefix}")
|
||||
|
||||
# Map resize method strings to PIL resize methods
|
||||
resize_methods = {
|
||||
"nearest-exact": Image.NEAREST,
|
||||
"bilinear": Image.BILINEAR,
|
||||
"bicubic": Image.BICUBIC,
|
||||
"lanczos": Image.LANCZOS,
|
||||
"box": Image.BOX
|
||||
}
|
||||
|
||||
# Handle different versions of PIL
|
||||
if hasattr(Image, 'Resampling'):
|
||||
resize_methods = {
|
||||
"nearest-exact": Image.Resampling.NEAREST,
|
||||
"bilinear": Image.Resampling.BILINEAR,
|
||||
"bicubic": Image.Resampling.BICUBIC,
|
||||
"lanczos": Image.Resampling.LANCZOS,
|
||||
"box": Image.Resampling.BOX
|
||||
}
|
||||
|
||||
# Get the selected resize method, default to LANCZOS if not found
|
||||
selected_resize_method = resize_methods.get(resize_method, Image.LANCZOS)
|
||||
|
||||
# Initialize Discord sender if enabled
|
||||
discord_success = False
|
||||
if send_to_discord and webhook_url:
|
||||
print(f"Discord integration enabled, preparing to send images to webhook")
|
||||
discord_success = True # Will be set to False if any send fails
|
||||
|
||||
# Initialize message_prefix for all Discord messages
|
||||
# This ensures prompts have a place to be attached regardless of other options
|
||||
image_info["message_prefix"] = ""
|
||||
|
||||
elif send_to_discord and not webhook_url:
|
||||
print("Discord integration was enabled but no webhook URL was provided")
|
||||
|
||||
# Add image info to Discord message if relevant options are enabled
|
||||
if send_to_discord and webhook_url and (add_date or add_time or add_dimensions or resize_to_power_of_2 or include_format_in_message):
|
||||
info_message = "\n\n**Image Information:**\n"
|
||||
|
||||
if "date" in image_info:
|
||||
info_message += f"**Date:** {image_info['date']}\n"
|
||||
|
||||
if "time" in image_info:
|
||||
info_message += f"**Time:** {image_info['time']}\n"
|
||||
|
||||
# Add format to the message if the option is enabled
|
||||
if include_format_in_message:
|
||||
info_message += f"**Format:** {file_format.upper()}\n"
|
||||
|
||||
# Update the message prefix with the information
|
||||
image_info["message_prefix"] = info_message
|
||||
|
||||
# Note: We don't add to discord_message yet, as dimensions aren't known until processing
|
||||
print("Prepared image information section for Discord message")
|
||||
|
||||
# Extract prompts if requested
|
||||
if send_to_discord and include_prompts_in_message:
|
||||
workflow_data = None
|
||||
|
||||
# First try to get workflow from extra_pnginfo
|
||||
if original_extra_pnginfo is not None and isinstance(original_extra_pnginfo, dict) and "workflow" in original_extra_pnginfo:
|
||||
workflow_data = original_extra_pnginfo["workflow"]
|
||||
|
||||
# If no workflow in extra_pnginfo, check if prompt is actually a workflow
|
||||
if workflow_data is None and original_prompt is not None:
|
||||
# Check if prompt is already a workflow
|
||||
if isinstance(original_prompt, dict) and "nodes" in original_prompt:
|
||||
workflow_data = original_prompt
|
||||
|
||||
# Extract prompts from workflow data
|
||||
if workflow_data is not None:
|
||||
positive_prompt, negative_prompt = extract_prompts_from_workflow(workflow_data)
|
||||
|
||||
# Ensure the prompts are strings or None
|
||||
if positive_prompt is not False and positive_prompt is not None and not isinstance(positive_prompt, str):
|
||||
positive_prompt = str(positive_prompt)
|
||||
print(f"Converted positive prompt to string: {positive_prompt[:50]}...")
|
||||
|
||||
if negative_prompt is not False and negative_prompt is not None and not isinstance(negative_prompt, str):
|
||||
negative_prompt = str(negative_prompt)
|
||||
print(f"Converted negative prompt to string: {negative_prompt[:50]}...")
|
||||
|
||||
# Check if we have valid prompt data
|
||||
has_valid_prompt = (
|
||||
(isinstance(positive_prompt, str) and positive_prompt) or
|
||||
(isinstance(negative_prompt, str) and negative_prompt)
|
||||
)
|
||||
|
||||
# Add prompts to Discord message if found
|
||||
if has_valid_prompt:
|
||||
prompt_message = "\n\n**Generation Prompts:**\n"
|
||||
|
||||
if isinstance(positive_prompt, str) and positive_prompt:
|
||||
prompt_message += f"**Positive:**\n```\n{positive_prompt}\n```\n"
|
||||
|
||||
if isinstance(negative_prompt, str) and negative_prompt:
|
||||
prompt_message += f"**Negative:**\n```\n{negative_prompt}\n```\n"
|
||||
|
||||
# Store prompt message for adding after image info
|
||||
image_info["prompt_message"] = prompt_message
|
||||
print("Prepared prompts for Discord message")
|
||||
|
||||
for batch_number, image in enumerate(images):
|
||||
# Convert the tensor to a PIL image
|
||||
# Optimization: Use torch operations for scaling/clipping/casting to avoid large float64 intermediate arrays on CPU
|
||||
# This is significantly faster (~70%) and uses less memory
|
||||
i = (image * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
img = Image.fromarray(i)
|
||||
|
||||
# Get original dimensions before any resizing
|
||||
orig_width, orig_height = img.size
|
||||
|
||||
# Resize to power of 2 if enabled
|
||||
if resize_to_power_of_2 == "enable":
|
||||
# Calculate nearest power of 2 dimensions
|
||||
new_width = 2 ** int(np.log2(orig_width) + 0.5) # Round to nearest power of 2
|
||||
new_height = 2 ** int(np.log2(orig_height) + 0.5) # Round to nearest power of 2
|
||||
|
||||
print(f"Resizing image from {orig_width}x{orig_height} to {new_width}x{new_height} (power of 2)")
|
||||
|
||||
# Store original and resized dimensions for Discord message
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
image_info["original_dimensions"] = f"{orig_width}x{orig_height}"
|
||||
image_info["resized_dimensions"] = f"{new_width}x{new_height}"
|
||||
|
||||
# Only resize if dimensions changed
|
||||
if (new_width != orig_width or new_height != orig_height):
|
||||
try:
|
||||
img = img.resize((new_width, new_height), selected_resize_method)
|
||||
print(f"Successfully resized using {resize_method} method")
|
||||
except Exception as e:
|
||||
print(f"Error during power of 2 resize: {e}")
|
||||
# Fallback to BICUBIC if selected method fails
|
||||
img = img.resize((new_width, new_height), Image.BICUBIC)
|
||||
print("Fallback to BICUBIC resize method due to error")
|
||||
|
||||
# Get dimensions - either original or resized
|
||||
width, height = img.size
|
||||
|
||||
# Add dimensions to filename if enabled
|
||||
dimensions_suffix = ""
|
||||
if add_dimensions == "enable":
|
||||
dimensions_suffix = f"_{width}x{height}"
|
||||
filename_prefix += dimensions_suffix
|
||||
|
||||
# Store dimensions for Discord message if needed
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
image_info["dimensions"] = f"{width}x{height}"
|
||||
|
||||
# Add image information to Discord message if this is the first image
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
# Add image info if available
|
||||
if "message_prefix" in image_info:
|
||||
info_message = image_info["message_prefix"]
|
||||
|
||||
# Add dimensions info if available
|
||||
if "original_dimensions" in image_info and "resized_dimensions" in image_info:
|
||||
info_message += f"**Original Dimensions:** {image_info['original_dimensions']}\n"
|
||||
info_message += f"**Resized Dimensions:** {image_info['resized_dimensions']} (Power of 2)\n"
|
||||
elif "dimensions" in image_info:
|
||||
info_message += f"**Dimensions:** {image_info['dimensions']}\n"
|
||||
|
||||
# Add the complete info message to the Discord message
|
||||
discord_message += info_message
|
||||
print("Added image information to Discord message")
|
||||
|
||||
# Add prompts after image info if available (decoupled from image info presence)
|
||||
if "prompt_message" in image_info:
|
||||
discord_message += image_info["prompt_message"]
|
||||
print("Added prompts to Discord message after image information")
|
||||
|
||||
# Create metadata for the image
|
||||
metadata = None
|
||||
if not args.disable_metadata:
|
||||
metadata = PngInfo()
|
||||
if prompt is not None:
|
||||
# Final sanitization check before embedding
|
||||
sanitized_prompt = sanitize_json_for_export(prompt)
|
||||
metadata.add_text("prompt", json.dumps(sanitized_prompt))
|
||||
if extra_pnginfo is not None:
|
||||
# Final sanitization check before embedding
|
||||
sanitized_extra_pnginfo = sanitize_json_for_export(extra_pnginfo)
|
||||
for x in sanitized_extra_pnginfo:
|
||||
if x == "workflow":
|
||||
# Extra sanitization for workflow data
|
||||
workflow_data = sanitize_json_for_export(sanitized_extra_pnginfo[x])
|
||||
metadata.add_text(x, json.dumps(workflow_data))
|
||||
else:
|
||||
metadata.add_text(x, json.dumps(sanitized_extra_pnginfo[x]))
|
||||
|
||||
# For Discord output
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(batch_number))
|
||||
|
||||
# Add dimensions tag before the counter if enabled
|
||||
if add_dimensions == "enable" and dimensions_suffix not in filename_with_batch_num:
|
||||
# Insert dimensions before counter
|
||||
base_name = os.path.splitext(filename_with_batch_num)[0]
|
||||
filename_with_batch_num = f"{base_name}{dimensions_suffix}"
|
||||
|
||||
# File extension based on format
|
||||
extension = f".{file_format}"
|
||||
file = f"{filename_with_batch_num}_{counter:05}_{extension}"
|
||||
|
||||
# Remove the additional underscore before the extension
|
||||
if file.endswith(f"_{extension}"):
|
||||
file = file[:-len(f"_{extension}")] + extension
|
||||
|
||||
filepath = os.path.join(full_output_folder, file)
|
||||
|
||||
try:
|
||||
# Save the image based on format
|
||||
if file_format == "png":
|
||||
# For PNG, make sure we have sanitized metadata
|
||||
if metadata is not None and hasattr(metadata, "text"):
|
||||
# Double check any JSON in the metadata
|
||||
for key in list(metadata.text.keys()):
|
||||
try:
|
||||
value = metadata.text[key]
|
||||
# Try to parse and sanitize any JSON values
|
||||
json_value = json.loads(value)
|
||||
sanitized_json = sanitize_json_for_export(json_value)
|
||||
metadata.text[key] = json.dumps(sanitized_json)
|
||||
except:
|
||||
# Not JSON or error, leave as is
|
||||
pass
|
||||
|
||||
img.save(filepath, pnginfo=metadata, compress_level=self.compress_level)
|
||||
elif file_format == "jpeg":
|
||||
# JPEG is always lossy, but we can set quality to maximum if lossless is requested
|
||||
jpeg_quality = 100 if lossless else quality
|
||||
img.save(filepath, format="JPEG", quality=jpeg_quality)
|
||||
elif file_format == "webp":
|
||||
if lossless:
|
||||
img.save(filepath, format="WEBP", lossless=True)
|
||||
else:
|
||||
img.save(filepath, format="WEBP", quality=quality)
|
||||
|
||||
output_files.append(filepath)
|
||||
|
||||
# Print dimensions for verification
|
||||
print(f"Saved image with dimensions: {img.size[0]}x{img.size[1]}")
|
||||
|
||||
# Add to results for UI display
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": "discord_output/" + (subfolder if subfolder else "") if save_output else "",
|
||||
"type": "output" if save_output else "temp",
|
||||
"path": filepath
|
||||
})
|
||||
|
||||
# Send to Discord if enabled
|
||||
if send_to_discord and webhook_url:
|
||||
try:
|
||||
# Generate unique filename for Discord using the selected format
|
||||
discord_filename = f"{uuid4()}.{file_format}"
|
||||
file_bytes = BytesIO()
|
||||
|
||||
# Optimization: Use PIL directly for JPEG/WebP to avoid numpy conversion overhead
|
||||
# Use CV2 for PNG as it is significantly faster for that format
|
||||
|
||||
# Optimization: Use PIL for JPEG encoding directly (faster, less memory)
|
||||
# Keep OpenCV for PNG (faster) and Pillow for WebP (legacy/consistency)
|
||||
|
||||
if file_format == "jpeg":
|
||||
# JPEG does not support RGBA, convert to RGB if needed
|
||||
save_img = img
|
||||
if save_img.mode == 'RGBA':
|
||||
save_img = save_img.convert('RGB')
|
||||
|
||||
jpeg_quality = 100 if lossless else quality
|
||||
save_img.save(file_bytes, format="JPEG", quality=jpeg_quality)
|
||||
file_bytes.seek(0)
|
||||
|
||||
elif file_format == "png":
|
||||
# Use CV2 for PNG (significantly faster)
|
||||
img_cv = np.array(img)
|
||||
|
||||
# Convert RGB (PIL) to BGR (OpenCV)
|
||||
if len(img_cv.shape) == 3 and img_cv.shape[2] == 3:
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGB2BGR)
|
||||
|
||||
# Handle color conversion for special cases
|
||||
if len(img_cv.shape) == 2: # Grayscale
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_GRAY2BGR)
|
||||
elif len(img_cv.shape) == 3 and img_cv.shape[2] == 4: # RGBA
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGBA2BGRA)
|
||||
|
||||
_, buffer = cv2.imencode('.png', img_cv)
|
||||
file_bytes = BytesIO(buffer)
|
||||
|
||||
elif file_format == "webp":
|
||||
try:
|
||||
if lossless:
|
||||
img.save(file_bytes, format="WEBP", lossless=True)
|
||||
else:
|
||||
img.save(file_bytes, format="WEBP", quality=quality)
|
||||
file_bytes.seek(0)
|
||||
except Exception as e:
|
||||
print(f"Error with WebP encoding for Discord: {e}, falling back to PNG")
|
||||
# Fallback to PNG if WebP encoding fails (using PIL)
|
||||
discord_filename = f"{os.path.splitext(discord_filename)[0]}.png"
|
||||
file_bytes = BytesIO() # Reset buffer
|
||||
img.save(file_bytes, format="PNG", compress_level=self.compress_level)
|
||||
file_bytes.seek(0)
|
||||
|
||||
# If batch grouping is enabled, store the files for later
|
||||
if group_batched_images:
|
||||
# Store this image for batch sending
|
||||
batch_discord_files.append((discord_filename, file_bytes.getvalue()))
|
||||
|
||||
# Prepare workflow JSON only once for the whole batch
|
||||
if batch_number == 0 and send_workflow_json and (prompt is not None or extra_pnginfo is not None):
|
||||
try:
|
||||
workflow_json = None
|
||||
|
||||
# First try to get workflow from extra_pnginfo
|
||||
if extra_pnginfo is not None and isinstance(extra_pnginfo, dict) and "workflow" in extra_pnginfo:
|
||||
workflow_json = extra_pnginfo["workflow"]
|
||||
|
||||
# If no workflow in extra_pnginfo, check if prompt is actually a workflow
|
||||
if workflow_json is None and prompt is not None:
|
||||
# Check if prompt is already a workflow
|
||||
if isinstance(prompt, dict) and "nodes" in prompt and "links" in prompt:
|
||||
workflow_json = prompt
|
||||
|
||||
# Ensure the workflow is sanitized
|
||||
if workflow_json:
|
||||
workflow_json = sanitize_json_for_export(workflow_json)
|
||||
batch_workflow_json = workflow_json
|
||||
except Exception as e:
|
||||
print(f"Error preparing workflow JSON for batch: {e}")
|
||||
|
||||
# Store discord message only once
|
||||
if batch_number == 0 and discord_message:
|
||||
batch_discord_data["content"] = discord_message
|
||||
|
||||
else:
|
||||
# Original non-batched behavior - send immediately
|
||||
# Prepare the Discord request
|
||||
files = {
|
||||
"file": (discord_filename, file_bytes.getvalue())
|
||||
}
|
||||
|
||||
# If enabled, also send the workflow JSON
|
||||
if send_workflow_json and (prompt is not None or extra_pnginfo is not None):
|
||||
try:
|
||||
workflow_json = None
|
||||
|
||||
# First check if extra_pnginfo contains the workflow data
|
||||
if extra_pnginfo is not None and "workflow" in extra_pnginfo:
|
||||
workflow_json = extra_pnginfo["workflow"]
|
||||
|
||||
# If no workflow in extra_pnginfo, check if prompt is actually a workflow
|
||||
if workflow_json is None and prompt is not None:
|
||||
# Check if prompt is already a workflow
|
||||
if isinstance(prompt, dict) and "nodes" in prompt and "links" in prompt:
|
||||
workflow_json = prompt
|
||||
|
||||
# Ensure the workflow is sanitized
|
||||
if workflow_json:
|
||||
workflow_json = sanitize_json_for_export(workflow_json)
|
||||
|
||||
# Generate a JSON file with the same base name
|
||||
json_filename = f"{os.path.splitext(discord_filename)[0]}.json"
|
||||
|
||||
# Convert workflow data to JSON string in the proper format
|
||||
json_data = json.dumps(workflow_json, indent=2)
|
||||
|
||||
# Add JSON file to the request
|
||||
files["workflow"] = (json_filename, json_data.encode('utf-8'))
|
||||
|
||||
print(f"ComfyUI workflow JSON file {json_filename} will be sent alongside the image")
|
||||
else:
|
||||
print("No workflow data found in the provided metadata")
|
||||
except Exception as e:
|
||||
print(f"Error preparing workflow JSON: {e}")
|
||||
|
||||
data = {}
|
||||
if discord_message:
|
||||
data["content"] = discord_message
|
||||
|
||||
# Only send to Discord if not batching
|
||||
if not group_batched_images:
|
||||
# Send to Discord with retry logic
|
||||
response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=files,
|
||||
data=data
|
||||
)
|
||||
|
||||
# Discord can return either 204 (no content) or 200 (success with content) for successful requests
|
||||
if response.status_code in [200, 204]:
|
||||
print(f"Successfully sent image {batch_number+1} to Discord")
|
||||
discord_sent_files.append(discord_filename)
|
||||
if send_workflow_json and "workflow" in files:
|
||||
print(f"Successfully sent workflow JSON for image {batch_number+1}")
|
||||
|
||||
# Try to extract CDN URLs from batch response
|
||||
if save_cdn_urls and response.status_code == 200:
|
||||
try:
|
||||
response_data = response.json()
|
||||
print(f"Received JSON response from Discord with {len(response_data) if isinstance(response_data, dict) else 'invalid'} fields")
|
||||
|
||||
if "attachments" in response_data and isinstance(response_data["attachments"], list):
|
||||
print(f"Found {len(response_data['attachments'])} attachments in Discord response")
|
||||
|
||||
for idx, attachment in enumerate(response_data["attachments"]):
|
||||
if "url" in attachment and "filename" in attachment:
|
||||
# Filter out workflow JSON files from URLs list
|
||||
if not attachment["filename"].endswith(".json"):
|
||||
batch_cdn_urls.append((attachment["filename"], attachment["url"]))
|
||||
print(f"Extracted CDN URL for batch image {idx+1}: {attachment['url']}")
|
||||
else:
|
||||
print(f"Skipping JSON file: {attachment['filename']}")
|
||||
else:
|
||||
print(f"Attachment {idx+1} missing URL or filename: {attachment.keys()}")
|
||||
|
||||
print(f"Total batch CDN URLs collected: {len(batch_cdn_urls)}")
|
||||
|
||||
# Create and send a text file with the CDN URLs if we have any
|
||||
if batch_cdn_urls:
|
||||
try:
|
||||
# Create the text file content
|
||||
url_text_content = "# Discord CDN URLs\n\n"
|
||||
for idx, (filename, url) in enumerate(batch_cdn_urls):
|
||||
url_text_content += f"{idx+1}. {filename}: {url}\n"
|
||||
|
||||
# Create a unique filename for the text file
|
||||
urls_filename = f"cdn_urls-{uuid4()}.txt"
|
||||
|
||||
# Prepare the request with just the URL file
|
||||
url_files = {"file": (urls_filename, url_text_content.encode('utf-8'))}
|
||||
url_data = {"content": "Discord CDN URLs for the uploaded images:"}
|
||||
|
||||
# Send a follow-up message with just the URLs text file
|
||||
url_response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=url_files,
|
||||
data=url_data
|
||||
)
|
||||
|
||||
if url_response.status_code in [200, 204]:
|
||||
print(f"Successfully sent CDN URLs text file to Discord")
|
||||
else:
|
||||
print(f"Error sending CDN URLs text file: Status code {url_response.status_code}")
|
||||
except Exception as e:
|
||||
print(f"Error creating or sending CDN URLs text file: {e}")
|
||||
except Exception as e:
|
||||
print(f"Error extracting CDN URLs from batch response: {e}")
|
||||
else:
|
||||
print(f"Error: Discord returned status code {response.status_code}")
|
||||
discord_send_success = False
|
||||
else:
|
||||
# Just mark it as queued for batch sending
|
||||
print(f"Image {batch_number+1} queued for batch sending to Discord")
|
||||
except Exception as e:
|
||||
print(f"Error processing image for Discord: {e}")
|
||||
discord_send_success = False
|
||||
|
||||
# Increment counter if not overwriting
|
||||
if overwrite_last != "enable":
|
||||
counter += 1
|
||||
except Exception as e:
|
||||
print(f"Error saving image: {e}")
|
||||
|
||||
if results:
|
||||
if save_output:
|
||||
print(f"DiscordSendSaveImage: Saved {len(results)} images to {full_output_folder}")
|
||||
else:
|
||||
print("DiscordSendSaveImage: Preview only mode - no images saved to disk")
|
||||
|
||||
# Discord status
|
||||
if send_to_discord and discord_sent_files:
|
||||
print("DiscordSendSaveImage: Successfully sent all images to Discord")
|
||||
|
||||
# If we have CDN URLs and we're not in batch mode, send them as a text file
|
||||
if save_cdn_urls and discord_cdn_urls and not (group_batched_images and len(images) > 1):
|
||||
try:
|
||||
# Create the text file content
|
||||
url_text_content = "# Discord CDN URLs\n\n"
|
||||
for idx, (filename, url) in enumerate(discord_cdn_urls):
|
||||
url_text_content += f"{idx+1}. {filename}: {url}\n"
|
||||
|
||||
# Create a unique filename for the text file
|
||||
urls_filename = f"cdn_urls-{uuid4()}.txt"
|
||||
|
||||
# Prepare the request with just the URL file
|
||||
url_files = {"file": (urls_filename, url_text_content.encode('utf-8'))}
|
||||
url_data = {"content": "Discord CDN URLs for the uploaded images:"}
|
||||
|
||||
# Send a follow-up message with just the URLs text file
|
||||
url_response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=url_files,
|
||||
data=url_data
|
||||
)
|
||||
|
||||
if url_response.status_code in [200, 204]:
|
||||
print(f"Successfully sent CDN URLs text file to Discord")
|
||||
else:
|
||||
print(f"Error sending CDN URLs text file: Status code {url_response.status_code}")
|
||||
except Exception as e:
|
||||
print(f"Error creating or sending CDN URLs text file: {e}")
|
||||
elif send_to_discord and not discord_send_success:
|
||||
print("DiscordSendSaveImage: There were errors sending some images to Discord")
|
||||
else:
|
||||
print("DiscordSendSaveImage: No images were processed")
|
||||
|
||||
# Send batch to Discord if enabled and we have images
|
||||
if send_to_discord and webhook_url and group_batched_images and batch_discord_files:
|
||||
try:
|
||||
print(f"Sending {len(batch_discord_files)} images as a batch to Discord...")
|
||||
|
||||
# Prepare files dictionary for the request
|
||||
files = {}
|
||||
for i, (filename, file_bytes) in enumerate(batch_discord_files):
|
||||
files[f"file{i}"] = (filename, file_bytes)
|
||||
|
||||
# Add workflow JSON if available
|
||||
if send_workflow_json and batch_workflow_json:
|
||||
try:
|
||||
# Generate a JSON file with a unique name
|
||||
json_filename = f"workflow-{uuid4()}.json"
|
||||
|
||||
# Convert workflow data to JSON string in the proper format
|
||||
json_data = json.dumps(batch_workflow_json, indent=2)
|
||||
|
||||
# Add JSON file to the request
|
||||
files["workflow"] = (json_filename, json_data.encode('utf-8'))
|
||||
|
||||
print(f"Adding workflow JSON file to batch Discord message")
|
||||
except Exception as e:
|
||||
print(f"Error preparing workflow JSON for batch: {e}")
|
||||
|
||||
# Send the batch to Discord with retry logic
|
||||
response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=files,
|
||||
data=batch_discord_data
|
||||
)
|
||||
|
||||
# Discord can return either 204 (no content) or 200 (success with content) for successful requests
|
||||
if response.status_code in [200, 204]:
|
||||
print(f"Successfully sent batch of {len(batch_discord_files)} images to Discord as a gallery")
|
||||
discord_send_success = True
|
||||
discord_sent_files = ["batch_gallery"] # Mark as successfully sent
|
||||
|
||||
# Try to extract CDN URLs from batch response
|
||||
if save_cdn_urls and response.status_code == 200:
|
||||
try:
|
||||
response_data = response.json()
|
||||
print(f"Received JSON response from Discord with {len(response_data) if isinstance(response_data, dict) else 'invalid'} fields")
|
||||
|
||||
if "attachments" in response_data and isinstance(response_data["attachments"], list):
|
||||
print(f"Found {len(response_data['attachments'])} attachments in Discord response")
|
||||
|
||||
for idx, attachment in enumerate(response_data["attachments"]):
|
||||
if "url" in attachment and "filename" in attachment:
|
||||
# Filter out workflow JSON files from URLs list
|
||||
if not attachment["filename"].endswith(".json"):
|
||||
batch_cdn_urls.append((attachment["filename"], attachment["url"]))
|
||||
print(f"Extracted CDN URL for batch image {idx+1}: {attachment['url']}")
|
||||
else:
|
||||
print(f"Skipping JSON file: {attachment['filename']}")
|
||||
else:
|
||||
print(f"Attachment {idx+1} missing URL or filename: {attachment.keys()}")
|
||||
|
||||
print(f"Total batch CDN URLs collected: {len(batch_cdn_urls)}")
|
||||
|
||||
# Create and send a text file with the CDN URLs if we have any
|
||||
if batch_cdn_urls:
|
||||
try:
|
||||
# Create the text file content
|
||||
url_text_content = "# Discord CDN URLs\n\n"
|
||||
for idx, (filename, url) in enumerate(batch_cdn_urls):
|
||||
url_text_content += f"{idx+1}. {filename}: {url}\n"
|
||||
|
||||
# Create a unique filename for the text file
|
||||
urls_filename = f"cdn_urls-{uuid4()}.txt"
|
||||
|
||||
# Prepare the request with just the URL file
|
||||
url_files = {"file": (urls_filename, url_text_content.encode('utf-8'))}
|
||||
url_data = {"content": "Discord CDN URLs for the uploaded images:"}
|
||||
|
||||
# Send a follow-up message with just the URLs text file
|
||||
url_response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=url_files,
|
||||
data=url_data
|
||||
)
|
||||
|
||||
if url_response.status_code in [200, 204]:
|
||||
print(f"Successfully sent CDN URLs text file to Discord")
|
||||
else:
|
||||
print(f"Error sending CDN URLs text file: Status code {url_response.status_code}")
|
||||
except Exception as e:
|
||||
print(f"Error creating or sending CDN URLs text file: {e}")
|
||||
except Exception as e:
|
||||
print(f"Error extracting CDN URLs from batch response: {e}")
|
||||
else:
|
||||
print(f"Error sending batch to Discord: Status code {response.status_code} - {response.text}")
|
||||
discord_send_success = False
|
||||
except Exception as e:
|
||||
print(f"Error sending batch to Discord: {e}")
|
||||
discord_send_success = False
|
||||
|
||||
# Update GitHub repository with CDN URLs if enabled - MOVED HERE AFTER ALL DISCORD OPERATIONS
|
||||
if github_cdn_update and send_to_discord and (discord_cdn_urls or batch_cdn_urls):
|
||||
# Use whichever list of URLs we have
|
||||
urls_to_send = discord_cdn_urls if discord_cdn_urls else batch_cdn_urls
|
||||
|
||||
print(f"GitHub update is enabled with: repo={github_repo}, token_provided={'Yes' if github_token else 'No'}, file_path={github_file_path}")
|
||||
print(f"Number of available CDN URLs to update GitHub: {len(urls_to_send)}")
|
||||
|
||||
if urls_to_send:
|
||||
# Call the GitHub update function
|
||||
print(f"Updating GitHub repository {github_repo} with {len(urls_to_send)} Discord CDN URLs...")
|
||||
success, message = update_github_cdn_urls(
|
||||
github_repo=github_repo,
|
||||
github_token=github_token,
|
||||
file_path=github_file_path,
|
||||
cdn_urls=urls_to_send
|
||||
)
|
||||
if success:
|
||||
print(f"GitHub update successful: {message}")
|
||||
else:
|
||||
print(f"GitHub update failed: {message}")
|
||||
else:
|
||||
print("No CDN URLs available to update GitHub repository")
|
||||
elif github_cdn_update:
|
||||
# If GitHub update is enabled but not triggered, explain why
|
||||
reasons = []
|
||||
if not send_to_discord:
|
||||
reasons.append("send_to_discord is disabled")
|
||||
if not (discord_cdn_urls or batch_cdn_urls):
|
||||
reasons.append("no CDN URLs were collected (did Discord upload succeed?)")
|
||||
if not github_repo:
|
||||
reasons.append("github_repo is empty")
|
||||
if not github_token:
|
||||
reasons.append("github_token is empty")
|
||||
if not github_file_path:
|
||||
reasons.append("github_file_path is empty")
|
||||
|
||||
print(f"GitHub update was enabled but not triggered because: {', '.join(reasons)}")
|
||||
|
||||
# Control UI preview based on show_preview flag
|
||||
if show_preview:
|
||||
return {"ui": {"images": results}, "result": ((save_output, output_files, discord_send_success if send_to_discord else None),)}, output_files[0] if output_files else ""
|
||||
else:
|
||||
# Return a minimal UI object without images
|
||||
return {"ui": {}, "result": ((save_output, output_files, discord_send_success if send_to_discord else None),)}, output_files[0] if output_files else ""
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, images, filename_prefix="ComfyUI-Image", overwrite_last=False,
|
||||
file_format="png", quality=95, lossless=True, add_date=False, add_time=False,
|
||||
add_dimensions=False, resize_to_power_of_2=False, save_output=True,
|
||||
resize_method="lanczos", show_preview=True, send_to_discord=False, webhook_url="", discord_message="",
|
||||
include_prompts_in_message=False, include_format_in_message=False, group_batched_images=True,
|
||||
send_workflow_json=False, save_cdn_urls=False, github_cdn_update=False, github_repo="",
|
||||
github_token="", github_file_path="cdn_urls.md", prompt=None, extra_pnginfo=None):
|
||||
return True
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,19 +0,0 @@
|
||||
"""
|
||||
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,11 @@
|
||||
"""
|
||||
ComfyUI-DiscordSend Node Implementations
|
||||
|
||||
This package contains the ComfyUI custom nodes for sending media to Discord.
|
||||
"""
|
||||
|
||||
from .base_node import BaseDiscordNode
|
||||
from .image_node import DiscordSendSaveImage
|
||||
from .video_node import DiscordSendSaveVideo
|
||||
|
||||
__all__ = ['BaseDiscordNode', 'DiscordSendSaveImage', 'DiscordSendSaveVideo']
|
||||
@@ -0,0 +1,343 @@
|
||||
"""
|
||||
Base class for Discord-enabled ComfyUI nodes.
|
||||
|
||||
Provides common functionality for sending media to Discord,
|
||||
including INPUT_TYPES definitions, sanitization, and Discord integration.
|
||||
"""
|
||||
|
||||
import os
|
||||
import folder_paths
|
||||
|
||||
from shared import (
|
||||
sanitize_json_for_export,
|
||||
update_github_cdn_urls,
|
||||
send_to_discord_with_retry,
|
||||
build_filename_with_metadata,
|
||||
get_output_directory,
|
||||
build_metadata_section,
|
||||
build_prompt_section,
|
||||
extract_cdn_urls_from_response,
|
||||
send_cdn_urls_file,
|
||||
extract_prompts_from_workflow
|
||||
)
|
||||
|
||||
|
||||
class BaseDiscordNode:
|
||||
"""
|
||||
Base class for Discord-enabled ComfyUI nodes.
|
||||
|
||||
Provides common functionality for:
|
||||
- Filename generation with metadata
|
||||
- Output directory management
|
||||
- Discord webhook integration
|
||||
- GitHub CDN URL updates
|
||||
- Workflow data sanitization
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.type = "output"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
self.output_dir = None
|
||||
|
||||
@staticmethod
|
||||
def get_discord_input_types():
|
||||
"""
|
||||
Returns Discord-related INPUT_TYPES fields.
|
||||
|
||||
These can be merged into a node's INPUT_TYPES definition.
|
||||
"""
|
||||
return {
|
||||
"send_to_discord": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "If enabled, will send the media to Discord via webhook."
|
||||
}),
|
||||
"webhook_url": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "Discord webhook URL. Get this from Discord server settings > Integrations > Webhooks."
|
||||
}),
|
||||
"discord_message": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Message to include with the media when sending to Discord."
|
||||
}),
|
||||
"include_prompts_in_message": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "If enabled, will include the generation prompts in the Discord message."
|
||||
}),
|
||||
"send_workflow_json": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "If enabled, will send the workflow JSON alongside the media."
|
||||
}),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def get_cdn_input_types():
|
||||
"""
|
||||
Returns CDN and GitHub-related INPUT_TYPES fields.
|
||||
"""
|
||||
return {
|
||||
"save_cdn_urls": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "If enabled, will extract and save Discord CDN URLs."
|
||||
}),
|
||||
"github_cdn_update": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "If enabled, will update a GitHub repository with the CDN URLs."
|
||||
}),
|
||||
"github_repo": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "GitHub repository in format 'username/repo'."
|
||||
}),
|
||||
"github_token": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "GitHub Personal Access Token (PAT). Generate at: Settings > Developer settings > Tokens. \n⚠️ Requires 'repo' scope (or 'public_repo') to upload files."
|
||||
}),
|
||||
"github_file_path": ("STRING", {
|
||||
"default": "cdn_urls.md",
|
||||
"multiline": False,
|
||||
"tooltip": "Path to the file in the GitHub repository to update."
|
||||
}),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def get_filename_input_types(
|
||||
add_date_default: bool = False,
|
||||
add_time_default: bool = True,
|
||||
add_dimensions_default: bool = False
|
||||
):
|
||||
"""
|
||||
Returns filename metadata INPUT_TYPES fields.
|
||||
|
||||
Args:
|
||||
add_date_default: Default value for add_date
|
||||
add_time_default: Default value for add_time
|
||||
add_dimensions_default: Default value for add_dimensions
|
||||
"""
|
||||
return {
|
||||
"add_date": ("BOOLEAN", {
|
||||
"default": add_date_default,
|
||||
"tooltip": "Add date (YYYY-MM-DD) to the filename."
|
||||
}),
|
||||
"add_time": ("BOOLEAN", {
|
||||
"default": add_time_default,
|
||||
"tooltip": "Add time (HH-MM-SS) to the filename."
|
||||
}),
|
||||
"add_dimensions": ("BOOLEAN", {
|
||||
"default": add_dimensions_default,
|
||||
"tooltip": "Add dimensions (WxH) to the filename."
|
||||
}),
|
||||
}
|
||||
|
||||
def sanitize_workflow_data(self, prompt, extra_pnginfo):
|
||||
"""
|
||||
Sanitize workflow data by removing sensitive information.
|
||||
|
||||
Args:
|
||||
prompt: The prompt data
|
||||
extra_pnginfo: Extra PNG info including workflow
|
||||
|
||||
Returns:
|
||||
Tuple of (sanitized_prompt, sanitized_extra_pnginfo,
|
||||
original_prompt, original_extra_pnginfo)
|
||||
"""
|
||||
# Store original references for prompt extraction
|
||||
original_prompt = prompt
|
||||
original_extra_pnginfo = extra_pnginfo
|
||||
|
||||
# Sanitize workflow data
|
||||
if prompt is not None:
|
||||
prompt = sanitize_json_for_export(prompt)
|
||||
|
||||
if extra_pnginfo is not None:
|
||||
extra_pnginfo = sanitize_json_for_export(extra_pnginfo)
|
||||
|
||||
return prompt, extra_pnginfo, original_prompt, original_extra_pnginfo
|
||||
|
||||
def build_filename_prefix(
|
||||
self,
|
||||
filename_prefix: str,
|
||||
add_date: bool,
|
||||
add_time: bool,
|
||||
add_dimensions: bool = False,
|
||||
width: int = None,
|
||||
height: int = None
|
||||
):
|
||||
"""
|
||||
Build filename prefix with metadata.
|
||||
|
||||
Returns:
|
||||
Tuple of (modified_prefix, info_dict)
|
||||
"""
|
||||
info_dict = {}
|
||||
filename_prefix, info_dict = build_filename_with_metadata(
|
||||
prefix=filename_prefix,
|
||||
add_date=add_date,
|
||||
add_time=add_time,
|
||||
add_dimensions=add_dimensions,
|
||||
width=width,
|
||||
height=height,
|
||||
info_dict=info_dict
|
||||
)
|
||||
filename_prefix += self.prefix_append
|
||||
return filename_prefix, info_dict
|
||||
|
||||
def get_dest_folder(self, save_output: bool):
|
||||
"""
|
||||
Get the destination folder for output files.
|
||||
|
||||
Args:
|
||||
save_output: Whether to save to output directory (True) or temp (False)
|
||||
|
||||
Returns:
|
||||
Path to the destination folder
|
||||
"""
|
||||
return get_output_directory(
|
||||
save_output=save_output,
|
||||
comfy_output_dir=folder_paths.get_output_directory(),
|
||||
temp_dir=folder_paths.get_temp_directory()
|
||||
)
|
||||
|
||||
def extract_workflow_from_metadata(self, original_prompt, original_extra_pnginfo):
|
||||
"""
|
||||
Extract workflow data from metadata.
|
||||
|
||||
Args:
|
||||
original_prompt: Original prompt data
|
||||
original_extra_pnginfo: Original extra PNG info
|
||||
|
||||
Returns:
|
||||
Workflow data dict, or None if not found
|
||||
"""
|
||||
workflow_data = None
|
||||
|
||||
# First try to get workflow from extra_pnginfo
|
||||
if (original_extra_pnginfo is not None and
|
||||
isinstance(original_extra_pnginfo, dict) and
|
||||
"workflow" in original_extra_pnginfo):
|
||||
workflow_data = original_extra_pnginfo["workflow"]
|
||||
|
||||
# If no workflow in extra_pnginfo, check if prompt is actually a workflow
|
||||
if workflow_data is None and original_prompt is not None:
|
||||
if isinstance(original_prompt, dict) and "nodes" in original_prompt:
|
||||
workflow_data = original_prompt
|
||||
|
||||
return workflow_data
|
||||
|
||||
def build_prompt_message(self, workflow_data):
|
||||
"""
|
||||
Extract and build prompt message from workflow data.
|
||||
|
||||
Args:
|
||||
workflow_data: Workflow data dict
|
||||
|
||||
Returns:
|
||||
Formatted prompt section string, or empty string
|
||||
"""
|
||||
if workflow_data is None:
|
||||
return ""
|
||||
|
||||
positive_prompt, negative_prompt = extract_prompts_from_workflow(workflow_data)
|
||||
return build_prompt_section(positive_prompt, negative_prompt)
|
||||
|
||||
def send_discord_files(
|
||||
self,
|
||||
webhook_url: str,
|
||||
files: dict,
|
||||
data: dict,
|
||||
save_cdn_urls: bool = False
|
||||
):
|
||||
"""
|
||||
Send files to Discord via webhook.
|
||||
|
||||
Args:
|
||||
webhook_url: Discord webhook URL
|
||||
files: Files dict for the request
|
||||
data: Data dict for the request
|
||||
save_cdn_urls: Whether to extract CDN URLs from response
|
||||
|
||||
Returns:
|
||||
Tuple of (success, response, cdn_urls)
|
||||
"""
|
||||
cdn_urls = []
|
||||
|
||||
response = send_to_discord_with_retry(
|
||||
webhook_url,
|
||||
files=files,
|
||||
data=data
|
||||
)
|
||||
|
||||
success = response.status_code in [200, 204]
|
||||
|
||||
if success and save_cdn_urls:
|
||||
cdn_urls = extract_cdn_urls_from_response(response)
|
||||
|
||||
return success, response, cdn_urls
|
||||
|
||||
def send_cdn_urls_to_discord(
|
||||
self,
|
||||
webhook_url: str,
|
||||
cdn_urls: list,
|
||||
message: str = "Discord CDN URLs:"
|
||||
):
|
||||
"""
|
||||
Send CDN URLs as a text file to Discord.
|
||||
|
||||
Args:
|
||||
webhook_url: Discord webhook URL
|
||||
cdn_urls: List of (filename, url) tuples
|
||||
message: Message to accompany the file
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
if not cdn_urls:
|
||||
return False
|
||||
|
||||
return send_cdn_urls_file(
|
||||
webhook_url=webhook_url,
|
||||
urls=cdn_urls,
|
||||
send_func=send_to_discord_with_retry,
|
||||
message=message
|
||||
)
|
||||
|
||||
def update_github_cdn(
|
||||
self,
|
||||
cdn_urls: list,
|
||||
github_repo: str,
|
||||
github_token: str,
|
||||
github_file_path: str
|
||||
):
|
||||
"""
|
||||
Update GitHub repository with CDN URLs.
|
||||
|
||||
Args:
|
||||
cdn_urls: List of (filename, url) tuples
|
||||
github_repo: GitHub repository in 'owner/repo' format
|
||||
github_token: GitHub personal access token
|
||||
github_file_path: Path to file in repository
|
||||
|
||||
Returns:
|
||||
Tuple of (success, message)
|
||||
"""
|
||||
if not cdn_urls:
|
||||
return False, "No CDN URLs to update"
|
||||
|
||||
print(f"Updating GitHub repository {github_repo} with {len(cdn_urls)} CDN URLs...")
|
||||
|
||||
success, message = update_github_cdn_urls(
|
||||
github_repo=github_repo,
|
||||
github_token=github_token,
|
||||
file_path=github_file_path,
|
||||
cdn_urls=cdn_urls
|
||||
)
|
||||
|
||||
if success:
|
||||
print(f"GitHub update successful: {message}")
|
||||
else:
|
||||
print(f"GitHub update failed: {message}")
|
||||
|
||||
return success, message
|
||||
@@ -0,0 +1,559 @@
|
||||
"""ComfyUI node for sending images to Discord and saving them locally."""
|
||||
|
||||
import os
|
||||
import json
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import folder_paths
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
from comfy.cli_args import args
|
||||
import re
|
||||
import cv2
|
||||
from io import BytesIO
|
||||
from uuid import uuid4
|
||||
from typing import Any, Union, List, Optional
|
||||
|
||||
# Import shared utilities
|
||||
from shared import (
|
||||
sanitize_token_from_text,
|
||||
process_batched_images,
|
||||
validate_path_is_safe,
|
||||
sanitize_json_for_export
|
||||
)
|
||||
|
||||
|
||||
from .base_node import BaseDiscordNode
|
||||
|
||||
|
||||
class DiscordSendSaveImage(BaseDiscordNode):
|
||||
"""
|
||||
A ComfyUI node that can send images to Discord and save them with advanced options.
|
||||
Images can be sent to Discord via webhook integration, while providing flexible
|
||||
saving options with customizable naming conventions and format options.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.type = "output"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
self.output_dir = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# Get base inputs from BaseDiscordNode
|
||||
base_inputs = BaseDiscordNode.get_discord_input_types()
|
||||
cdn_inputs = BaseDiscordNode.get_cdn_input_types()
|
||||
filename_inputs = BaseDiscordNode.get_filename_input_types(add_date_default=False)
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE", {"tooltip": "The images to save and/or send to Discord."}),
|
||||
"filename_prefix": ("STRING", {"default": "ComfyUI-Image", "tooltip": "The prefix for the saved files. Supports %batch_num% placeholder for batch indexing."}),
|
||||
"overwrite_last": ("BOOLEAN", {"default": False, "tooltip": "⚠️ CAUTION: If enabled, new saves will REPLACE the previous file with the same name. Useful for iterative testing, dangerous for batch production. Note: You must also disable 'add_time' and 'add_date' to ensure filenames are identical."})
|
||||
},
|
||||
"optional": {
|
||||
"file_format": (["png", "jpeg", "webp"], {
|
||||
"default": "png",
|
||||
"tooltip": "The format to save images in. PNG is lossless but larger. JPEG and WebP are smaller but lossy."
|
||||
}),
|
||||
"quality": ("INT", {
|
||||
"default": 95,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"step": 1,
|
||||
"tooltip": "Quality (1-100) for JPEG/WebP. Ignored for PNG. Higher values = better quality but larger file size."
|
||||
}),
|
||||
"lossless": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Use lossless compression for WebP (PNG is always lossless). For JPEG, forces maximum quality (100)."
|
||||
}),
|
||||
"save_output": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Whether to save images to disk. When disabled, images will only be previewed in the UI."
|
||||
}),
|
||||
"show_preview": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Whether to show image previews in the UI. Disable to reduce UI clutter for large batches."
|
||||
}),
|
||||
"resize_to_power_of_2": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Resize images to nearest power of 2 dimensions (useful for game textures). ⚠️ May distort aspect ratio. Uses the algorithm selected in 'resize_method'."
|
||||
}),
|
||||
"resize_method": (["nearest-exact", "bilinear", "bicubic", "lanczos", "box"], {
|
||||
"default": "lanczos",
|
||||
"tooltip": "Resampling algorithm used ONLY when 'resize_to_power_of_2' is enabled. Ignored otherwise. \n• lanczos: Best for photos\n• nearest-exact: Best for pixel art\n• bilinear/bicubic: Faster"
|
||||
}),
|
||||
"include_format_in_message": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Whether to include the image format in the Discord message."
|
||||
}),
|
||||
"group_batched_images": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Group all images from a batch into a single Discord message with a gallery, rather than sending each one separately. Maximum is 9 images."
|
||||
}),
|
||||
# Mix in shared options
|
||||
**filename_inputs,
|
||||
**base_inputs,
|
||||
**cdn_inputs,
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO"
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("image_path",)
|
||||
FUNCTION = "save_images"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "image/output"
|
||||
DESCRIPTION = "Saves images with advanced options and can send them to Discord via webhook integration. Returns the path to the first saved image."
|
||||
|
||||
@classmethod
|
||||
def CONTEXT_MENUS(s):
|
||||
return {
|
||||
"Show Preview": lambda self, **kwargs: {"show_preview": True},
|
||||
"Hide Preview": lambda self, **kwargs: {"show_preview": False},
|
||||
}
|
||||
|
||||
def save_images(self, images, filename_prefix="ComfyUI-Image", overwrite_last=False,
|
||||
file_format="png", quality=95, lossless=True, add_date=False, add_time=False,
|
||||
add_dimensions=False, resize_to_power_of_2=False, save_output=True,
|
||||
resize_method="lanczos", show_preview=True, send_to_discord=False, webhook_url="", discord_message="",
|
||||
include_prompts_in_message=False, include_format_in_message=False, send_workflow_json=False,
|
||||
group_batched_images=True, save_cdn_urls=False, github_cdn_update=False, github_repo="",
|
||||
github_token="", github_file_path="cdn_urls.md", prompt=None, extra_pnginfo=None):
|
||||
"""
|
||||
Save images for and optionally send to Discord.
|
||||
"""
|
||||
results = []
|
||||
output_files = []
|
||||
discord_sent_files = []
|
||||
discord_send_success = True
|
||||
|
||||
# For batch grouping
|
||||
batch_discord_files = []
|
||||
batch_discord_data = {}
|
||||
batch_workflow_json = None
|
||||
|
||||
# For tracking Discord CDN URLs
|
||||
discord_cdn_urls = []
|
||||
batch_cdn_urls = []
|
||||
|
||||
# 1. Sanitize workflow data using base class method
|
||||
prompt, extra_pnginfo, original_prompt, original_extra_pnginfo = self.sanitize_workflow_data(
|
||||
prompt, extra_pnginfo
|
||||
)
|
||||
|
||||
# 2. Build filename prefix with metadata using base class method
|
||||
filename_prefix, image_info = self.build_filename_prefix(
|
||||
filename_prefix, add_date, add_time, False, None, None
|
||||
)
|
||||
|
||||
# Add prefix append
|
||||
filename_prefix += self.prefix_append
|
||||
|
||||
# 3. Get output directory using base class method
|
||||
dest_folder = self.get_dest_folder(save_output)
|
||||
|
||||
# Setup paths using ComfyUI's path validation
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix, dest_folder, images[0].shape[1], images[0].shape[0])
|
||||
|
||||
# For overwrite functionality, we'll just always use the same counter instead of bypassing validation
|
||||
if overwrite_last:
|
||||
counter = 1 # Always use the same counter value for overwriting
|
||||
else:
|
||||
# When not overwriting, we need to find the highest existing counter and start from there
|
||||
# This ensures we're always creating new files
|
||||
try:
|
||||
# Get all existing files with this prefix
|
||||
base_filename = os.path.basename(filename).replace("%batch_num%", "")
|
||||
existing_files = [f for f in os.listdir(full_output_folder)
|
||||
if os.path.basename(f).startswith(base_filename)]
|
||||
|
||||
if existing_files:
|
||||
# Extract counters from filenames
|
||||
existing_counters = []
|
||||
for f in existing_files:
|
||||
# Extract counter pattern (5 digits) from filename
|
||||
counter_match = re.search(r'_(\d{5})\.', f)
|
||||
if counter_match:
|
||||
existing_counters.append(int(counter_match.group(1)))
|
||||
|
||||
# Also try alternative pattern where the counter is followed by extension
|
||||
counter_match = re.search(r'_(\d{5})_\.', f)
|
||||
if counter_match:
|
||||
existing_counters.append(int(counter_match.group(1)))
|
||||
|
||||
# Set counter to one more than the highest existing counter
|
||||
if existing_counters:
|
||||
counter = max(existing_counters) + 1
|
||||
except Exception as e:
|
||||
print(f"Error determining next file counter: {e}")
|
||||
# Default to ComfyUI's counter if we can't determine the next one
|
||||
|
||||
print(f"Using counter: {counter} for {'overwriting' if overwrite_last else 'new files'}")
|
||||
print(f"Output prefix: {filename_prefix}")
|
||||
|
||||
# Map resize method strings to PIL resize methods
|
||||
resize_methods = {
|
||||
"nearest-exact": Image.NEAREST,
|
||||
"bilinear": Image.BILINEAR,
|
||||
"bicubic": Image.BICUBIC,
|
||||
"lanczos": Image.LANCZOS,
|
||||
"box": Image.BOX
|
||||
}
|
||||
|
||||
# Handle different versions of PIL
|
||||
if hasattr(Image, 'Resampling'):
|
||||
resize_methods = {
|
||||
"nearest-exact": Image.Resampling.NEAREST,
|
||||
"bilinear": Image.Resampling.BILINEAR,
|
||||
"bicubic": Image.Resampling.BICUBIC,
|
||||
"lanczos": Image.Resampling.LANCZOS,
|
||||
"box": Image.Resampling.BOX
|
||||
}
|
||||
|
||||
# Get the selected resize method, default to LANCZOS if not found
|
||||
selected_resize_method = resize_methods.get(resize_method, Image.LANCZOS)
|
||||
|
||||
# Initialize Discord sender if enabled
|
||||
discord_success = False
|
||||
if send_to_discord and webhook_url:
|
||||
print(f"Discord integration enabled, preparing to send images to webhook")
|
||||
discord_success = True # Will be set to False if any send fails
|
||||
|
||||
# Initialize message_prefix for all Discord messages
|
||||
# This ensures prompts have a place to be attached regardless of other options
|
||||
image_info["message_prefix"] = ""
|
||||
|
||||
elif send_to_discord and not webhook_url:
|
||||
print("Discord integration was enabled but no webhook URL was provided")
|
||||
|
||||
# Build image info message using shared utility
|
||||
if send_to_discord and webhook_url and (add_date or add_time or add_dimensions or resize_to_power_of_2 or include_format_in_message):
|
||||
info_message = build_metadata_section(
|
||||
info_dict=image_info,
|
||||
include_date=add_date,
|
||||
include_time=add_time,
|
||||
include_dimensions=False, # Dimensions added later after processing
|
||||
include_format=include_format_in_message,
|
||||
file_format=file_format,
|
||||
section_title="Image Information"
|
||||
)
|
||||
image_info["message_prefix"] = info_message
|
||||
print("Prepared image information section for Discord message")
|
||||
|
||||
# 4. Extract and build prompts section
|
||||
if send_to_discord and include_prompts_in_message:
|
||||
workflow_data = self.extract_workflow_from_metadata(original_prompt, original_extra_pnginfo)
|
||||
if workflow_data:
|
||||
prompt_message = self.build_prompt_message(workflow_data)
|
||||
if prompt_message:
|
||||
image_info["prompt_message"] = prompt_message
|
||||
print("Prepared prompts for Discord message")
|
||||
|
||||
# Optimization: Create metadata once for the entire batch
|
||||
# This prevents redundant sanitization and JSON serialization for every image
|
||||
metadata = None
|
||||
if not args.disable_metadata:
|
||||
metadata = PngInfo()
|
||||
if prompt is not None:
|
||||
# Prompt is already sanitized at start of function
|
||||
metadata.add_text("prompt", json.dumps(prompt))
|
||||
if extra_pnginfo is not None:
|
||||
# extra_pnginfo is already sanitized at start of function
|
||||
for x in extra_pnginfo:
|
||||
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
|
||||
|
||||
batch_counter = 0
|
||||
for chunk in process_batched_images(images):
|
||||
if len(chunk.shape) == 4:
|
||||
chunk_images = [chunk[i] for i in range(chunk.shape[0])]
|
||||
else:
|
||||
chunk_images = [chunk]
|
||||
|
||||
for image_np in chunk_images:
|
||||
batch_number = batch_counter
|
||||
batch_counter += 1
|
||||
# Convert the tensor to a PIL image
|
||||
i = image_np
|
||||
img = Image.fromarray(i)
|
||||
|
||||
# Track if resizing happened to optimize Discord encoding later
|
||||
was_resized = False
|
||||
orig_width, orig_height = img.size
|
||||
|
||||
# Resize to power of 2 if enabled
|
||||
if resize_to_power_of_2:
|
||||
new_width = 2 ** int(np.log2(orig_width) + 0.5)
|
||||
new_height = 2 ** int(np.log2(orig_height) + 0.5)
|
||||
|
||||
print(f"Resizing image from {orig_width}x{orig_height} to {new_width}x{new_height} (power of 2)")
|
||||
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
image_info["original_dimensions"] = f"{orig_width}x{orig_height}"
|
||||
image_info["resized_dimensions"] = f"{new_width}x{new_height}"
|
||||
|
||||
if (new_width != orig_width or new_height != orig_height):
|
||||
try:
|
||||
img = img.resize((new_width, new_height), selected_resize_method)
|
||||
was_resized = True
|
||||
print(f"Successfully resized using {resize_method} method")
|
||||
except Exception as e:
|
||||
print(f"Error during power of 2 resize: {e}")
|
||||
img = img.resize((new_width, new_height), Image.BICUBIC)
|
||||
was_resized = True
|
||||
print("Fallback to BICUBIC resize method due to error")
|
||||
|
||||
# Get dimensions
|
||||
width, height = img.size
|
||||
|
||||
# Add dimensions to filename if enabled
|
||||
dimensions_suffix = ""
|
||||
if add_dimensions:
|
||||
dimensions_suffix = f"_{width}x{height}"
|
||||
filename_prefix += dimensions_suffix
|
||||
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
image_info["dimensions"] = f"{width}x{height}"
|
||||
|
||||
# Add image information to Discord message if this is the first image
|
||||
if send_to_discord and webhook_url and batch_number == 0:
|
||||
if "message_prefix" in image_info:
|
||||
info_message = image_info["message_prefix"]
|
||||
has_resize_dimensions = "original_dimensions" in image_info and "resized_dimensions" in image_info
|
||||
has_dimensions = "dimensions" in image_info
|
||||
|
||||
if (has_resize_dimensions or has_dimensions) and not info_message:
|
||||
info_message = "\n\n**Image Information:**\n"
|
||||
|
||||
if has_resize_dimensions:
|
||||
info_message += f"**Original Dimensions:** {image_info['original_dimensions']}\n"
|
||||
info_message += f"**Resized Dimensions:** {image_info['resized_dimensions']} (Power of 2)\n"
|
||||
elif has_dimensions:
|
||||
info_message += f"**Dimensions:** {image_info['dimensions']}\n"
|
||||
|
||||
if info_message:
|
||||
discord_message += info_message
|
||||
print("Added image information to Discord message")
|
||||
|
||||
if "prompt_message" in image_info:
|
||||
discord_message += image_info["prompt_message"]
|
||||
print("Added prompts to Discord message after image information")
|
||||
|
||||
# For Discord output
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(batch_number))
|
||||
|
||||
# Add dimensions tag before the counter if enabled
|
||||
if add_dimensions and dimensions_suffix not in filename_with_batch_num:
|
||||
base_name = os.path.splitext(filename_with_batch_num)[0]
|
||||
filename_with_batch_num = f"{base_name}{dimensions_suffix}"
|
||||
|
||||
# File extension based on format
|
||||
extension = f".{file_format}"
|
||||
file = f"{filename_with_batch_num}_{counter:05}_{extension}"
|
||||
|
||||
if file.endswith(f"_{extension}"):
|
||||
file = file[:-len(f"_{extension}")] + extension
|
||||
|
||||
filepath = os.path.join(full_output_folder, file)
|
||||
|
||||
# Security: Validate output path to prevent symlink overwrites
|
||||
validate_path_is_safe(filepath, base_dir=full_output_folder)
|
||||
|
||||
try:
|
||||
# Save the image based on format
|
||||
if file_format == "png":
|
||||
img.save(filepath, pnginfo=metadata, compress_level=self.compress_level)
|
||||
elif file_format == "jpeg":
|
||||
jpeg_quality = 100 if lossless else quality
|
||||
img.save(filepath, format="JPEG", quality=jpeg_quality)
|
||||
elif file_format == "webp":
|
||||
if lossless:
|
||||
img.save(filepath, format="WEBP", lossless=True)
|
||||
else:
|
||||
img.save(filepath, format="WEBP", quality=quality)
|
||||
|
||||
output_files.append(filepath)
|
||||
|
||||
print(f"Saved image with dimensions: {img.size[0]}x{img.size[1]}")
|
||||
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": "discord_output/" + (subfolder if subfolder else "") if save_output else "",
|
||||
"type": "output" if save_output else "temp",
|
||||
"path": filepath
|
||||
})
|
||||
|
||||
# Send to Discord if enabled
|
||||
if send_to_discord and webhook_url:
|
||||
try:
|
||||
discord_filename = f"{uuid4()}.{file_format}"
|
||||
file_bytes = BytesIO()
|
||||
|
||||
if file_format == "jpeg":
|
||||
save_img = img
|
||||
if save_img.mode == 'RGBA':
|
||||
save_img = save_img.convert('RGB')
|
||||
jpeg_quality = 100 if lossless else quality
|
||||
save_img.save(file_bytes, format="JPEG", quality=jpeg_quality)
|
||||
file_bytes.seek(0)
|
||||
|
||||
elif file_format == "png":
|
||||
if not was_resized:
|
||||
img_cv = i
|
||||
else:
|
||||
img_cv = np.array(img)
|
||||
|
||||
if len(img_cv.shape) == 3 and img_cv.shape[2] == 3:
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGB2BGR)
|
||||
if len(img_cv.shape) == 2:
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_GRAY2BGR)
|
||||
elif len(img_cv.shape) == 3 and img_cv.shape[2] == 4:
|
||||
img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGBA2BGRA)
|
||||
|
||||
_, buffer = cv2.imencode('.png', img_cv)
|
||||
file_bytes = BytesIO(buffer)
|
||||
|
||||
elif file_format == "webp":
|
||||
try:
|
||||
if lossless:
|
||||
img.save(file_bytes, format="WEBP", lossless=True)
|
||||
else:
|
||||
img.save(file_bytes, format="WEBP", quality=quality)
|
||||
file_bytes.seek(0)
|
||||
except Exception as e:
|
||||
print(f"Error with WebP encoding for Discord: {e}, falling back to PNG")
|
||||
discord_filename = f"{os.path.splitext(discord_filename)[0]}.png"
|
||||
file_bytes = BytesIO() # Reset buffer
|
||||
img.save(file_bytes, format="PNG", compress_level=self.compress_level)
|
||||
file_bytes.seek(0)
|
||||
|
||||
if group_batched_images:
|
||||
batch_discord_files.append((discord_filename, file_bytes.getvalue()))
|
||||
|
||||
# Prepare workflow JSON only once for the whole batch
|
||||
if batch_number == 0 and send_workflow_json and (prompt is not None or extra_pnginfo is not None):
|
||||
wflow = self.extract_workflow_from_metadata(original_prompt, original_extra_pnginfo)
|
||||
if wflow:
|
||||
batch_workflow_json = wflow
|
||||
|
||||
if batch_number == 0 and discord_message:
|
||||
batch_discord_data["content"] = discord_message
|
||||
|
||||
else:
|
||||
# Immediate send
|
||||
files = {
|
||||
"file": (discord_filename, file_bytes.getvalue())
|
||||
}
|
||||
|
||||
if send_workflow_json and (prompt is not None or extra_pnginfo is not None):
|
||||
wflow = self.extract_workflow_from_metadata(original_prompt, original_extra_pnginfo)
|
||||
if wflow:
|
||||
# Sanitize to remove webhook URLs and tokens
|
||||
wflow = sanitize_json_for_export(wflow)
|
||||
json_filename = f"{os.path.splitext(discord_filename)[0]}.json"
|
||||
files["workflow"] = (json_filename, json.dumps(wflow, indent=2).encode('utf-8'))
|
||||
|
||||
data = {}
|
||||
if discord_message:
|
||||
data["content"] = discord_message
|
||||
|
||||
success, response, new_urls = self.send_discord_files(webhook_url, files, data, save_cdn_urls)
|
||||
|
||||
if success:
|
||||
print(f"Successfully sent image {batch_number+1} to Discord")
|
||||
discord_sent_files.append(discord_filename)
|
||||
if new_urls:
|
||||
batch_cdn_urls.extend(new_urls)
|
||||
self.send_cdn_urls_to_discord(webhook_url, new_urls, "Discord CDN URLs for the uploaded images:")
|
||||
else:
|
||||
print(f"Error: Discord returned status code {response.status_code}")
|
||||
discord_send_success = False
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing image for Discord: {e}")
|
||||
discord_send_success = False
|
||||
|
||||
if not overwrite_last:
|
||||
counter += 1
|
||||
except Exception as e:
|
||||
print(f"Error saving image: {e}")
|
||||
|
||||
if results:
|
||||
if save_output:
|
||||
print(f"DiscordSendSaveImage: Saved {len(results)} images to {full_output_folder}")
|
||||
else:
|
||||
print("DiscordSendSaveImage: Preview only mode - no images saved to disk")
|
||||
|
||||
if send_to_discord and discord_sent_files:
|
||||
print("DiscordSendSaveImage: Successfully sent all images to Discord")
|
||||
elif send_to_discord and not discord_send_success:
|
||||
print("DiscordSendSaveImage: There were errors sending some images to Discord")
|
||||
else:
|
||||
print("DiscordSendSaveImage: No images were processed")
|
||||
|
||||
# Send batch to Discord
|
||||
if send_to_discord and webhook_url and group_batched_images and batch_discord_files:
|
||||
try:
|
||||
print(f"Sending {len(batch_discord_files)} images as a batch to Discord...")
|
||||
|
||||
files = {}
|
||||
for i, (filename, file_bytes) in enumerate(batch_discord_files):
|
||||
files[f"file{i}"] = (filename, file_bytes)
|
||||
|
||||
if send_workflow_json and batch_workflow_json:
|
||||
# Sanitize to remove webhook URLs and tokens
|
||||
sanitized_workflow = sanitize_json_for_export(batch_workflow_json)
|
||||
json_filename = f"workflow-{uuid4()}.json"
|
||||
json_data = json.dumps(sanitized_workflow, indent=2)
|
||||
files["workflow"] = (json_filename, json_data.encode('utf-8'))
|
||||
|
||||
success, response, new_urls = self.send_discord_files(webhook_url, files, batch_discord_data, save_cdn_urls)
|
||||
|
||||
if success:
|
||||
print(f"Successfully sent batch of {len(batch_discord_files)} images to Discord as a gallery")
|
||||
discord_sent_files = ["batch_gallery"]
|
||||
if save_cdn_urls and new_urls:
|
||||
batch_cdn_urls.extend(new_urls)
|
||||
self.send_cdn_urls_to_discord(webhook_url, new_urls, "Discord CDN URLs for the uploaded images:")
|
||||
else:
|
||||
error_msg = sanitize_token_from_text(response.text, webhook_url)
|
||||
print(f"Error sending batch to Discord: Status code {response.status_code} - {error_msg}")
|
||||
discord_send_success = False
|
||||
except Exception as e:
|
||||
print(f"Error sending batch to Discord: {e}")
|
||||
discord_send_success = False
|
||||
|
||||
# Update GitHub repository
|
||||
if github_cdn_update and send_to_discord and (discord_cdn_urls or batch_cdn_urls):
|
||||
urls_to_send = discord_cdn_urls if discord_cdn_urls else batch_cdn_urls
|
||||
self.update_github_cdn(urls_to_send, github_repo, github_token, github_file_path)
|
||||
|
||||
elif github_cdn_update:
|
||||
reasons = []
|
||||
if not send_to_discord: reasons.append("send_to_discord is disabled")
|
||||
if not (discord_cdn_urls or batch_cdn_urls): reasons.append("no CDN URLs were collected")
|
||||
if not github_repo: reasons.append("github_repo is empty")
|
||||
if not github_token: reasons.append("github_token is empty")
|
||||
if not github_file_path: reasons.append("github_file_path is empty")
|
||||
print(f"GitHub update was enabled but not triggered because: {', '.join(reasons)}")
|
||||
|
||||
# Return results
|
||||
if show_preview:
|
||||
return {"ui": {"images": results}, "result": ((save_output, output_files, discord_send_success if send_to_discord else None),)}, output_files[0] if output_files else ""
|
||||
else:
|
||||
return {"ui": {}, "result": ((save_output, output_files, discord_send_success if send_to_discord else None),)}, output_files[0] if output_files else ""
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, images, filename_prefix="ComfyUI-Image", overwrite_last=False,
|
||||
file_format="png", quality=95, lossless=True, add_date=False, add_time=False,
|
||||
add_dimensions=False, resize_to_power_of_2=False, save_output=True,
|
||||
resize_method="lanczos", show_preview=True, send_to_discord=False, webhook_url="", discord_message="",
|
||||
include_prompts_in_message=False, include_format_in_message=False, group_batched_images=True,
|
||||
send_workflow_json=False, save_cdn_urls=False, github_cdn_update=False, github_repo="",
|
||||
github_token="", github_file_path="cdn_urls.md", prompt=None, extra_pnginfo=None):
|
||||
return True
|
||||
+1069
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-discordsend"
|
||||
description = "A ComfyUI extension that enables seamless sharing of AI-generated images and videos directly to Discord."
|
||||
version = "1.1.0"
|
||||
version = "2.0.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["requests>=2.25.0"]
|
||||
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
# Discord Bot Dependencies
|
||||
# Install with: pip install -r requirements-bot.txt
|
||||
#
|
||||
# Note: This includes node dependencies plus bot-specific packages
|
||||
|
||||
# Shared with nodes
|
||||
requests>=2.25.0
|
||||
Pillow>=9.0.0
|
||||
numpy>=1.20.0
|
||||
|
||||
# Discord bot framework
|
||||
discord.py>=2.3.0
|
||||
|
||||
# Async HTTP client for ComfyUI API
|
||||
aiohttp>=3.9.0
|
||||
|
||||
# Database ORM for job tracking
|
||||
sqlalchemy>=2.0.0
|
||||
|
||||
# Async SQLite driver
|
||||
aiosqlite>=0.19.0
|
||||
|
||||
# YAML config parsing
|
||||
pyyaml>=6.0.0
|
||||
@@ -0,0 +1,6 @@
|
||||
# Node-only Dependencies (same as requirements.txt)
|
||||
# This file exists for clarity - use requirements.txt for ComfyUI Manager
|
||||
|
||||
requests>=2.25.0
|
||||
Pillow>=9.0.0
|
||||
numpy>=1.20.0
|
||||
+9
-6
@@ -1,7 +1,10 @@
|
||||
# ComfyUI-DiscordSend Node Dependencies
|
||||
#
|
||||
# This file contains minimal dependencies for the ComfyUI nodes only.
|
||||
# ComfyUI Manager will automatically install these.
|
||||
#
|
||||
# For Discord bot support, manually run: pip install -r requirements-bot.txt
|
||||
|
||||
requests>=2.25.0
|
||||
discord.py>=2.3.0
|
||||
aiohttp>=3.9.0
|
||||
sqlalchemy>=2.0.0
|
||||
aiosqlite>=0.19.0
|
||||
pydantic>=2.0.0
|
||||
pyyaml>=6.0.0
|
||||
Pillow>=9.0.0
|
||||
numpy>=1.20.0
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
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,
|
||||
sanitize_token_from_text
|
||||
)
|
||||
from .discord.message_builder import (
|
||||
build_metadata_section,
|
||||
build_prompt_section,
|
||||
build_discord_message,
|
||||
format_file_size
|
||||
)
|
||||
from .discord.cdn_extractor import (
|
||||
extract_cdn_urls_from_response,
|
||||
send_cdn_urls_file
|
||||
)
|
||||
from .media.image_processing import tensor_to_numpy_uint8, process_batched_images
|
||||
from .github_integration import update_github_cdn_urls
|
||||
from .logging_config import setup_logging, get_logger
|
||||
from .filename_utils import build_filename_with_metadata, get_timestamp_string
|
||||
from .path_utils import get_output_directory, ensure_directory_exists, validate_path_is_safe
|
||||
|
||||
__all__ = [
|
||||
# Workflow utilities
|
||||
'sanitize_json_for_export',
|
||||
'extract_prompts_from_workflow',
|
||||
'WorkflowBuilder',
|
||||
# Discord utilities
|
||||
'DiscordWebhookClient',
|
||||
'validate_webhook_url',
|
||||
'send_to_discord_with_retry',
|
||||
'sanitize_token_from_text',
|
||||
# Discord message building
|
||||
'build_metadata_section',
|
||||
'build_prompt_section',
|
||||
'build_discord_message',
|
||||
'format_file_size',
|
||||
# CDN extraction
|
||||
'extract_cdn_urls_from_response',
|
||||
'send_cdn_urls_file',
|
||||
# Media utilities
|
||||
'tensor_to_numpy_uint8',
|
||||
'process_batched_images',
|
||||
# GitHub integration
|
||||
'update_github_cdn_urls',
|
||||
# Logging
|
||||
'setup_logging',
|
||||
'get_logger',
|
||||
# Filename utilities
|
||||
'build_filename_with_metadata',
|
||||
'get_timestamp_string',
|
||||
# Path utilities
|
||||
'get_output_directory',
|
||||
'ensure_directory_exists',
|
||||
'validate_path_is_safe',
|
||||
]
|
||||
@@ -0,0 +1,48 @@
|
||||
"""
|
||||
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
|
||||
)
|
||||
from .message_builder import (
|
||||
build_metadata_section,
|
||||
build_prompt_section,
|
||||
build_discord_message,
|
||||
validate_message_content,
|
||||
format_file_info,
|
||||
format_file_size
|
||||
)
|
||||
from .cdn_extractor import (
|
||||
extract_cdn_urls_from_response,
|
||||
create_cdn_urls_content,
|
||||
send_cdn_urls_file,
|
||||
collect_and_send_cdn_urls
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Webhook client
|
||||
'DiscordWebhookClient',
|
||||
'validate_webhook_url',
|
||||
'sanitize_webhook_for_logging',
|
||||
'send_to_discord_with_retry',
|
||||
'validate_file_for_discord',
|
||||
# Message building
|
||||
'build_metadata_section',
|
||||
'build_prompt_section',
|
||||
'build_discord_message',
|
||||
'validate_message_content',
|
||||
'format_file_info',
|
||||
'format_file_size',
|
||||
# CDN extraction
|
||||
'extract_cdn_urls_from_response',
|
||||
'create_cdn_urls_content',
|
||||
'send_cdn_urls_file',
|
||||
'collect_and_send_cdn_urls',
|
||||
]
|
||||
@@ -0,0 +1,168 @@
|
||||
"""
|
||||
Discord CDN URL extraction utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides functions for extracting CDN URLs from Discord responses
|
||||
and creating/sending URL text files.
|
||||
"""
|
||||
|
||||
from typing import List, Tuple, Optional, Any
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
def extract_cdn_urls_from_response(
|
||||
response: Any,
|
||||
exclude_json: bool = True
|
||||
) -> List[Tuple[str, str]]:
|
||||
"""
|
||||
Extract CDN URLs from a Discord webhook response.
|
||||
|
||||
Args:
|
||||
response: Response object from Discord API (must have .status_code and .json())
|
||||
exclude_json: Whether to exclude .json files from results
|
||||
|
||||
Returns:
|
||||
List of (filename, url) tuples
|
||||
"""
|
||||
cdn_urls = []
|
||||
|
||||
if response.status_code != 200:
|
||||
return cdn_urls
|
||||
|
||||
try:
|
||||
response_data = response.json()
|
||||
print(f"Received JSON response from Discord with "
|
||||
f"{len(response_data) if isinstance(response_data, dict) else 'invalid'} fields")
|
||||
|
||||
if "attachments" in response_data and isinstance(response_data["attachments"], list):
|
||||
print(f"Found {len(response_data['attachments'])} attachments in Discord response")
|
||||
|
||||
for idx, attachment in enumerate(response_data["attachments"]):
|
||||
if "url" in attachment and "filename" in attachment:
|
||||
filename = attachment["filename"]
|
||||
url = attachment["url"]
|
||||
|
||||
# Filter out workflow JSON files if requested
|
||||
if exclude_json and filename.endswith(".json"):
|
||||
print(f"Skipping JSON file: {filename}")
|
||||
continue
|
||||
|
||||
cdn_urls.append((filename, url))
|
||||
print(f"Extracted CDN URL for attachment {idx + 1}: {url}")
|
||||
else:
|
||||
print(f"Attachment {idx + 1} missing URL or filename: {attachment.keys()}")
|
||||
|
||||
print(f"Total CDN URLs collected: {len(cdn_urls)}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error extracting CDN URLs from response: {e}")
|
||||
|
||||
return cdn_urls
|
||||
|
||||
|
||||
def create_cdn_urls_content(
|
||||
urls: List[Tuple[str, str]],
|
||||
header: str = "# Discord CDN URLs\n\n"
|
||||
) -> str:
|
||||
"""
|
||||
Create text content from a list of CDN URLs.
|
||||
|
||||
Args:
|
||||
urls: List of (filename, url) tuples
|
||||
header: Header text for the content
|
||||
|
||||
Returns:
|
||||
Formatted text content
|
||||
"""
|
||||
content = header
|
||||
for idx, (filename, url) in enumerate(urls):
|
||||
content += f"{idx + 1}. {filename}: {url}\n"
|
||||
return content
|
||||
|
||||
|
||||
def send_cdn_urls_file(
|
||||
webhook_url: str,
|
||||
urls: List[Tuple[str, str]],
|
||||
send_func: Any,
|
||||
message: str = "Discord CDN URLs for the uploaded files:",
|
||||
filename_prefix: str = "cdn_urls"
|
||||
) -> bool:
|
||||
"""
|
||||
Create and send a text file containing CDN URLs to Discord.
|
||||
|
||||
Args:
|
||||
webhook_url: Discord webhook URL
|
||||
urls: List of (filename, url) tuples
|
||||
send_func: Function to send to Discord (send_to_discord_with_retry)
|
||||
message: Message to accompany the file
|
||||
filename_prefix: Prefix for the generated filename
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
"""
|
||||
if not urls:
|
||||
print("No CDN URLs to send")
|
||||
return False
|
||||
|
||||
try:
|
||||
# Create the text file content
|
||||
url_text_content = create_cdn_urls_content(urls)
|
||||
|
||||
# Create a unique filename for the text file
|
||||
urls_filename = f"{filename_prefix}-{uuid4()}.txt"
|
||||
|
||||
# Prepare the request with just the URL file
|
||||
url_files = {"file": (urls_filename, url_text_content.encode('utf-8'))}
|
||||
url_data = {"content": message}
|
||||
|
||||
# Send a follow-up message with just the URLs text file
|
||||
url_response = send_func(
|
||||
webhook_url,
|
||||
files=url_files,
|
||||
data=url_data
|
||||
)
|
||||
|
||||
if url_response.status_code in [200, 204]:
|
||||
print(f"Successfully sent CDN URLs text file to Discord")
|
||||
return True
|
||||
else:
|
||||
print(f"Error sending CDN URLs text file: Status code {url_response.status_code}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error creating or sending CDN URLs text file: {e}")
|
||||
return False
|
||||
|
||||
|
||||
def collect_and_send_cdn_urls(
|
||||
response: Any,
|
||||
webhook_url: str,
|
||||
send_func: Any,
|
||||
save_cdn_urls: bool,
|
||||
existing_urls: Optional[List[Tuple[str, str]]] = None,
|
||||
message: str = "Discord CDN URLs for the uploaded files:"
|
||||
) -> List[Tuple[str, str]]:
|
||||
"""
|
||||
Convenience function to extract CDN URLs from a response and optionally send them.
|
||||
|
||||
Args:
|
||||
response: Discord webhook response
|
||||
webhook_url: Webhook URL for sending the URLs file
|
||||
send_func: Function to send to Discord
|
||||
save_cdn_urls: Whether to extract and save CDN URLs
|
||||
existing_urls: Existing URLs to append to (for batch operations)
|
||||
message: Message to accompany the URLs file
|
||||
|
||||
Returns:
|
||||
List of all collected CDN URLs
|
||||
"""
|
||||
if existing_urls is None:
|
||||
existing_urls = []
|
||||
|
||||
if not save_cdn_urls:
|
||||
return existing_urls
|
||||
|
||||
# Extract URLs from this response
|
||||
new_urls = extract_cdn_urls_from_response(response)
|
||||
all_urls = existing_urls + new_urls
|
||||
|
||||
return all_urls
|
||||
@@ -0,0 +1,210 @@
|
||||
"""
|
||||
Discord message building utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides functions for constructing Discord messages with metadata,
|
||||
prompts, and other formatted content.
|
||||
"""
|
||||
|
||||
from typing import Dict, Any, Optional, Tuple, List
|
||||
|
||||
|
||||
def build_metadata_section(
|
||||
info_dict: Dict[str, Any],
|
||||
include_date: bool = True,
|
||||
include_time: bool = True,
|
||||
include_dimensions: bool = True,
|
||||
include_format: bool = True,
|
||||
file_format: Optional[str] = None,
|
||||
frame_rate: Optional[float] = None,
|
||||
section_title: str = "Information"
|
||||
) -> str:
|
||||
"""
|
||||
Build a formatted metadata section for Discord messages.
|
||||
|
||||
Args:
|
||||
info_dict: Dictionary containing metadata (date, time, dimensions, etc.)
|
||||
include_date: Whether to include date if present
|
||||
include_time: Whether to include time if present
|
||||
include_dimensions: Whether to include dimensions if present
|
||||
include_format: Whether to include format information
|
||||
file_format: File format string (e.g., "png", "mp4")
|
||||
frame_rate: Frame rate for video (optional)
|
||||
section_title: Title for the section (e.g., "Image Information", "Video Info")
|
||||
|
||||
Returns:
|
||||
Formatted metadata string, or empty string if no metadata
|
||||
"""
|
||||
metadata_lines = []
|
||||
|
||||
if include_date and "date" in info_dict:
|
||||
metadata_lines.append(f"**Date:** {info_dict['date']}")
|
||||
|
||||
if include_time and "time" in info_dict:
|
||||
metadata_lines.append(f"**Time:** {info_dict['time']}")
|
||||
|
||||
if include_dimensions and "dimensions" in info_dict:
|
||||
metadata_lines.append(f"**Dimensions:** {info_dict['dimensions']}")
|
||||
|
||||
if frame_rate is not None:
|
||||
metadata_lines.append(f"**Frame Rate:** {frame_rate} fps")
|
||||
|
||||
if include_format and file_format:
|
||||
metadata_lines.append(f"**Format:** {file_format.upper()}")
|
||||
|
||||
if not metadata_lines:
|
||||
return ""
|
||||
|
||||
section = f"\n\n**{section_title}:**\n"
|
||||
section += "\n".join(metadata_lines) + "\n"
|
||||
return section
|
||||
|
||||
|
||||
def build_prompt_section(
|
||||
positive_prompt: Optional[str],
|
||||
negative_prompt: Optional[str],
|
||||
section_title: str = "Generation Prompts"
|
||||
) -> str:
|
||||
"""
|
||||
Build a formatted prompts section for Discord messages.
|
||||
|
||||
Args:
|
||||
positive_prompt: The positive/main prompt text
|
||||
negative_prompt: The negative prompt text
|
||||
section_title: Title for the section
|
||||
|
||||
Returns:
|
||||
Formatted prompt string, or empty string if no prompts
|
||||
"""
|
||||
# Validate and normalize prompts
|
||||
if positive_prompt is not None and not isinstance(positive_prompt, str):
|
||||
positive_prompt = str(positive_prompt)
|
||||
if negative_prompt is not None and not isinstance(negative_prompt, str):
|
||||
negative_prompt = str(negative_prompt)
|
||||
|
||||
has_positive = isinstance(positive_prompt, str) and positive_prompt.strip()
|
||||
has_negative = isinstance(negative_prompt, str) and negative_prompt.strip()
|
||||
|
||||
if not has_positive and not has_negative:
|
||||
return ""
|
||||
|
||||
section = f"\n\n**{section_title}:**\n"
|
||||
|
||||
if has_positive:
|
||||
section += f"**Positive:**\n```\n{positive_prompt.strip()}\n```\n"
|
||||
|
||||
if has_negative:
|
||||
section += f"**Negative:**\n```\n{negative_prompt.strip()}\n```\n"
|
||||
|
||||
return section
|
||||
|
||||
|
||||
def build_discord_message(
|
||||
base_message: str = "",
|
||||
metadata_section: str = "",
|
||||
prompt_section: str = "",
|
||||
additional_sections: Optional[List[str]] = None,
|
||||
max_length: int = 2000
|
||||
) -> str:
|
||||
"""
|
||||
Build a complete Discord message from components.
|
||||
|
||||
Args:
|
||||
base_message: The main message content
|
||||
metadata_section: Pre-built metadata section
|
||||
prompt_section: Pre-built prompt section
|
||||
additional_sections: List of additional section strings
|
||||
max_length: Maximum message length (Discord limit is 2000)
|
||||
|
||||
Returns:
|
||||
Complete formatted message, truncated if necessary
|
||||
"""
|
||||
parts = [base_message] if base_message else []
|
||||
|
||||
if metadata_section:
|
||||
parts.append(metadata_section)
|
||||
|
||||
if prompt_section:
|
||||
parts.append(prompt_section)
|
||||
|
||||
if additional_sections:
|
||||
parts.extend(additional_sections)
|
||||
|
||||
message = "".join(parts)
|
||||
|
||||
# Truncate if necessary
|
||||
if len(message) > max_length:
|
||||
truncation_notice = "\n...[Message truncated]"
|
||||
message = message[:max_length - len(truncation_notice)] + truncation_notice
|
||||
|
||||
return message
|
||||
|
||||
|
||||
def validate_message_content(message: str) -> Tuple[bool, str]:
|
||||
"""
|
||||
Validate Discord message content.
|
||||
|
||||
Args:
|
||||
message: Message content to validate
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, validation_message)
|
||||
"""
|
||||
if not message:
|
||||
return True, "Empty message (valid for file-only uploads)"
|
||||
|
||||
if len(message) > 2000:
|
||||
return False, f"Message exceeds 2000 character limit ({len(message)} chars)"
|
||||
|
||||
# Check for required sections (informational)
|
||||
has_prompts = "Generation Prompts" in message
|
||||
|
||||
info_parts = []
|
||||
info_parts.append(f"Message has {message.count(chr(10))} lines")
|
||||
|
||||
if not has_prompts:
|
||||
info_parts.append("WARNING: Message does NOT contain 'Generation Prompts' section")
|
||||
|
||||
return True, "\n".join(info_parts)
|
||||
|
||||
|
||||
def format_file_info(
|
||||
filename: str,
|
||||
file_size: int,
|
||||
mime_type: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
Format file information for logging/display.
|
||||
|
||||
Args:
|
||||
filename: Name of the file
|
||||
file_size: Size in bytes
|
||||
mime_type: MIME type of the file
|
||||
|
||||
Returns:
|
||||
Formatted string with file information
|
||||
"""
|
||||
size_str = format_file_size(file_size)
|
||||
info = f"File: {filename} ({size_str})"
|
||||
if mime_type:
|
||||
info += f" [{mime_type}]"
|
||||
return info
|
||||
|
||||
|
||||
def format_file_size(size_bytes: int) -> str:
|
||||
"""
|
||||
Format file size in human-readable format.
|
||||
|
||||
Args:
|
||||
size_bytes: Size in bytes
|
||||
|
||||
Returns:
|
||||
Formatted string (e.g., "1.5 MB", "256 KB")
|
||||
"""
|
||||
if size_bytes < 1024:
|
||||
return f"{size_bytes} bytes"
|
||||
elif size_bytes < 1024 * 1024:
|
||||
return f"{size_bytes / 1024:.1f} KB"
|
||||
elif size_bytes < 1024 * 1024 * 1024:
|
||||
return f"{size_bytes / (1024 * 1024):.1f} MB"
|
||||
else:
|
||||
return f"{size_bytes / (1024 * 1024 * 1024):.2f} GB"
|
||||
@@ -20,8 +20,8 @@ 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-]+$",
|
||||
r"https://(?:www\.)?discord(?:app)?\.com/api/webhooks/\d+/[\w-]+$",
|
||||
r"https://(?:www\.)?discordapp\.com/api/webhooks/\d+/[\w-]+$",
|
||||
]
|
||||
|
||||
|
||||
@@ -38,8 +38,8 @@ def validate_webhook_url(url: str) -> Tuple[bool, str]:
|
||||
if not url:
|
||||
return False, "Webhook URL is empty"
|
||||
|
||||
if not url.startswith("http"):
|
||||
return False, "Webhook URL must start with http:// or https://"
|
||||
if not url.startswith("https://"):
|
||||
return False, "Webhook URL must start with https://"
|
||||
|
||||
# Check against known patterns
|
||||
for pattern in WEBHOOK_URL_PATTERNS:
|
||||
@@ -70,6 +70,31 @@ def sanitize_webhook_for_logging(url: str) -> str:
|
||||
return "[REDACTED_WEBHOOK_URL]"
|
||||
|
||||
|
||||
def sanitize_token_from_text(text: str, webhook_url: str) -> str:
|
||||
"""
|
||||
Sanitize the webhook token from arbitrary text.
|
||||
|
||||
Args:
|
||||
text: The text to sanitize
|
||||
webhook_url: The webhook URL containing the token
|
||||
|
||||
Returns:
|
||||
Text with the token replaced by [REDACTED]
|
||||
"""
|
||||
if not text or not webhook_url:
|
||||
return text
|
||||
|
||||
# Pattern: https://discord.com/api/webhooks/{id}/{token}
|
||||
# Use case-insensitive matching to handle potential variations
|
||||
match = re.search(r"/api/webhooks/\d+/([\w-]+)", webhook_url, re.IGNORECASE)
|
||||
if match:
|
||||
token = match.group(1)
|
||||
if token in text:
|
||||
return text.replace(token, "[REDACTED]")
|
||||
|
||||
return text
|
||||
|
||||
|
||||
class DiscordWebhookClient:
|
||||
"""
|
||||
Client for sending messages and files to Discord via webhooks.
|
||||
@@ -243,9 +268,10 @@ class DiscordWebhookClient:
|
||||
|
||||
# Client errors (don't retry)
|
||||
if 400 <= response.status_code < 500:
|
||||
sanitized_details = sanitize_token_from_text(response.text[:500], self.webhook_url)
|
||||
return False, {
|
||||
"error": f"Discord API error: {response.status_code}",
|
||||
"details": response.text[:500]
|
||||
"details": sanitized_details
|
||||
}
|
||||
|
||||
# Server errors (retry)
|
||||
@@ -254,7 +280,9 @@ class DiscordWebhookClient:
|
||||
except requests.exceptions.Timeout:
|
||||
last_error = "Request timed out"
|
||||
except requests.exceptions.RequestException as e:
|
||||
last_error = str(e)
|
||||
# Sanitize error message to prevent token leakage
|
||||
error_msg = sanitize_token_from_text(str(e), self.webhook_url)
|
||||
last_error = error_msg
|
||||
|
||||
# Exponential backoff
|
||||
if attempt < self.max_retries - 1:
|
||||
@@ -395,8 +423,27 @@ def send_to_discord_with_retry(
|
||||
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
|
||||
# Sanitize error message to prevent token leakage
|
||||
error_msg = sanitize_token_from_text(str(e), webhook_url)
|
||||
logger.warning(f"Request error: {error_msg}, attempt {attempt + 1}/{max_retries}")
|
||||
|
||||
# Store sanitized exception to avoid leaking token if raised later
|
||||
# Use case-insensitive matching to handle uppercase URLs
|
||||
match = re.search(r"/api/webhooks/\d+/([\w-]+)", webhook_url, re.IGNORECASE)
|
||||
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:
|
||||
@@ -0,0 +1,83 @@
|
||||
"""
|
||||
Filename utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides functions for building filenames with date, time, and dimension metadata.
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Dict, Optional, Tuple, Any
|
||||
|
||||
|
||||
def build_filename_with_metadata(
|
||||
prefix: str,
|
||||
add_date: bool = False,
|
||||
add_time: bool = False,
|
||||
add_dimensions: bool = False,
|
||||
width: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
info_dict: Optional[Dict[str, Any]] = None
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""
|
||||
Build a filename with optional date, time, and dimension suffixes.
|
||||
|
||||
Args:
|
||||
prefix: The base filename prefix
|
||||
add_date: Whether to add the current date (YYYY-MM-DD)
|
||||
add_time: Whether to add the current time (HH-MM-SS)
|
||||
add_dimensions: Whether to add dimensions (WxH)
|
||||
width: Image/video width (required if add_dimensions is True)
|
||||
height: Image/video height (required if add_dimensions is True)
|
||||
info_dict: Optional dict to update with metadata (creates new if None)
|
||||
|
||||
Returns:
|
||||
Tuple of (modified_prefix, info_dict with metadata)
|
||||
"""
|
||||
if info_dict is None:
|
||||
info_dict = {}
|
||||
|
||||
metadata_parts = []
|
||||
|
||||
if add_date:
|
||||
current_date = time.strftime("%Y-%m-%d")
|
||||
metadata_parts.append(current_date)
|
||||
info_dict["date"] = current_date
|
||||
print(f"Adding date to filename: {current_date}")
|
||||
|
||||
if add_time:
|
||||
current_time = time.strftime("%H-%M-%S")
|
||||
metadata_parts.append(current_time)
|
||||
info_dict["time"] = current_time
|
||||
print(f"Adding time to filename: {current_time}")
|
||||
|
||||
if add_dimensions and width is not None and height is not None:
|
||||
dim_text = f"{width}x{height}"
|
||||
metadata_parts.append(dim_text)
|
||||
info_dict["dimensions"] = dim_text
|
||||
print(f"Adding dimensions to filename: {dim_text}")
|
||||
|
||||
modified_prefix = prefix
|
||||
if metadata_parts:
|
||||
metadata_suffix = "_" + "_".join(metadata_parts)
|
||||
modified_prefix += metadata_suffix
|
||||
print(f"Final metadata suffix: {metadata_suffix}")
|
||||
|
||||
return modified_prefix, info_dict
|
||||
|
||||
|
||||
def get_timestamp_string(include_date: bool = True, include_time: bool = True) -> str:
|
||||
"""
|
||||
Get a formatted timestamp string.
|
||||
|
||||
Args:
|
||||
include_date: Include date in format YYYY-MM-DD
|
||||
include_time: Include time in format HH-MM-SS
|
||||
|
||||
Returns:
|
||||
Formatted timestamp string
|
||||
"""
|
||||
parts = []
|
||||
if include_date:
|
||||
parts.append(time.strftime("%Y-%m-%d"))
|
||||
if include_time:
|
||||
parts.append(time.strftime("%H-%M-%S"))
|
||||
return "_".join(parts) if parts else ""
|
||||
@@ -6,11 +6,43 @@ 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,
|
||||
@@ -44,9 +76,13 @@ def update_github_cdn_urls(
|
||||
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"
|
||||
# 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}"
|
||||
@@ -74,7 +110,11 @@ def update_github_cdn_urls(
|
||||
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}"
|
||||
# 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")
|
||||
@@ -120,7 +160,11 @@ def update_github_cdn_urls(
|
||||
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}"
|
||||
# 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"
|
||||
@@ -8,6 +8,22 @@ import logging
|
||||
import sys
|
||||
|
||||
|
||||
def setup_logging(level: int = logging.INFO) -> None:
|
||||
"""
|
||||
Set up logging configuration for the application.
|
||||
|
||||
Args:
|
||||
level: The logging level to use (default: INFO)
|
||||
"""
|
||||
logging.basicConfig(
|
||||
level=level,
|
||||
format='[%(name)s] %(levelname)s: %(message)s',
|
||||
handlers=[
|
||||
logging.StreamHandler(sys.stdout)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def get_logger(name: str = "comfyui_discordsend") -> logging.Logger:
|
||||
"""
|
||||
Get a configured logger for the extension.
|
||||
@@ -0,0 +1,40 @@
|
||||
"""
|
||||
Media Processing Utilities
|
||||
|
||||
Provides image and video processing functions.
|
||||
"""
|
||||
|
||||
from .image_processing import tensor_to_numpy_uint8
|
||||
from .format_utils import (
|
||||
parse_format_string,
|
||||
normalize_video_extension,
|
||||
get_mime_type,
|
||||
validate_video_for_discord,
|
||||
is_animated_format,
|
||||
supports_alpha
|
||||
)
|
||||
from .video_encoder import (
|
||||
detect_ffmpeg,
|
||||
FFmpegEncoder,
|
||||
PILEncoder,
|
||||
optimize_video_for_discord,
|
||||
mux_audio_to_video
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Image processing
|
||||
'tensor_to_numpy_uint8',
|
||||
# Format utilities
|
||||
'parse_format_string',
|
||||
'normalize_video_extension',
|
||||
'get_mime_type',
|
||||
'validate_video_for_discord',
|
||||
'is_animated_format',
|
||||
'supports_alpha',
|
||||
# Video encoding
|
||||
'detect_ffmpeg',
|
||||
'FFmpegEncoder',
|
||||
'PILEncoder',
|
||||
'optimize_video_for_discord',
|
||||
'mux_audio_to_video',
|
||||
]
|
||||
@@ -0,0 +1,137 @@
|
||||
"""
|
||||
Video format utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides format detection, extension mapping, and validation.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Tuple, Optional
|
||||
|
||||
|
||||
def parse_format_string(format_str: str) -> Tuple[str, str]:
|
||||
"""
|
||||
Parse a format string into type and extension.
|
||||
|
||||
Args:
|
||||
format_str: Format string like "video/h264-mp4" or "image/gif"
|
||||
|
||||
Returns:
|
||||
Tuple of (format_type, format_extension)
|
||||
"""
|
||||
if "/" in format_str:
|
||||
format_type, format_ext = format_str.split("/", 1)
|
||||
else:
|
||||
format_type = "video"
|
||||
format_ext = format_str
|
||||
|
||||
return format_type, format_ext
|
||||
|
||||
|
||||
def normalize_video_extension(format_str: str) -> str:
|
||||
"""
|
||||
Normalize a format string to a file extension.
|
||||
|
||||
Args:
|
||||
format_str: Format string like "video/h264-mp4"
|
||||
|
||||
Returns:
|
||||
Normalized extension (e.g., "mp4", "webm", "gif")
|
||||
"""
|
||||
_, format_ext = parse_format_string(format_str)
|
||||
|
||||
# Map format strings to extensions
|
||||
extension_map = {
|
||||
"h264-mp4": "mp4",
|
||||
"h265-mp4": "mp4",
|
||||
"vp9-webm": "webm",
|
||||
"prores": "mov",
|
||||
}
|
||||
|
||||
return extension_map.get(format_ext, format_ext)
|
||||
|
||||
|
||||
def get_mime_type(extension: str) -> str:
|
||||
"""
|
||||
Get MIME type for a video extension.
|
||||
|
||||
Args:
|
||||
extension: File extension (without dot)
|
||||
|
||||
Returns:
|
||||
MIME type string
|
||||
"""
|
||||
mime_types = {
|
||||
"mp4": "video/mp4",
|
||||
"webm": "video/webm",
|
||||
"gif": "image/gif",
|
||||
"mov": "video/quicktime",
|
||||
"avi": "video/x-msvideo",
|
||||
"mkv": "video/x-matroska",
|
||||
}
|
||||
return mime_types.get(extension.lower(), "application/octet-stream")
|
||||
|
||||
|
||||
def validate_video_for_discord(file_path: str, max_size_mb: int = 25) -> Tuple[bool, str]:
|
||||
"""
|
||||
Validate that a video file is compatible with Discord.
|
||||
|
||||
Args:
|
||||
file_path: Path to the video file
|
||||
max_size_mb: Maximum file size in megabytes (default 25MB for Discord)
|
||||
|
||||
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 = max_size_mb * 1024 * 1024
|
||||
if file_size > max_size:
|
||||
return False, f"File exceeds Discord's size limit of {max_size_mb}MB ({file_size / (1024*1024):.2f}MB)"
|
||||
|
||||
ext = os.path.splitext(file_path)[1].lower().lstrip('.')
|
||||
|
||||
if ext in ['mp4', 'webm', 'gif']:
|
||||
return True, "Valid"
|
||||
elif ext in ['mov']:
|
||||
return False, "MOV files may need conversion for Discord compatibility"
|
||||
elif ext in ['png', 'apng']:
|
||||
return False, "PNG/APNG sequence may need compilation for Discord"
|
||||
else:
|
||||
return False, f"Unknown format '{ext}' - may not be compatible with Discord"
|
||||
|
||||
|
||||
def is_animated_format(extension: str) -> bool:
|
||||
"""
|
||||
Check if a format supports animation.
|
||||
|
||||
Args:
|
||||
extension: File extension (without dot)
|
||||
|
||||
Returns:
|
||||
True if the format supports animation
|
||||
"""
|
||||
animated_formats = {'gif', 'webp', 'mp4', 'webm', 'mov', 'avi', 'mkv', 'apng'}
|
||||
return extension.lower() in animated_formats
|
||||
|
||||
|
||||
def supports_alpha(extension: str) -> bool:
|
||||
"""
|
||||
Check if a format supports alpha channel (transparency).
|
||||
|
||||
Args:
|
||||
extension: File extension (without dot)
|
||||
|
||||
Returns:
|
||||
True if the format supports alpha
|
||||
"""
|
||||
alpha_formats = {'webm', 'gif', 'webp', 'png', 'apng', 'mov'}
|
||||
return extension.lower() in alpha_formats
|
||||
@@ -0,0 +1,51 @@
|
||||
"""
|
||||
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)
|
||||
# Further Optimization: Use clamp_ (in-place) to avoid allocating a second float tensor
|
||||
return (tensor * 255.0).clamp_(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
|
||||
def process_batched_images(image_sequence, batch_size=20):
|
||||
"""
|
||||
Generator that processes images in batches to optimize GPU-CPU transfer.
|
||||
|
||||
Args:
|
||||
image_sequence: A torch.Tensor or list of tensors/images
|
||||
batch_size: Number of frames to process at once for Tensor inputs
|
||||
|
||||
Yields:
|
||||
Numpy array for each batch or frame, contiguous and ready for ffmpeg
|
||||
"""
|
||||
# Optimized path for Tensor input
|
||||
if isinstance(image_sequence, torch.Tensor):
|
||||
total = len(image_sequence)
|
||||
for i in range(0, total, batch_size):
|
||||
# Process a chunk of frames on GPU/CPU together
|
||||
# This amortizes the overhead of kernel launches and synchronization
|
||||
batch = image_sequence[i:i+batch_size]
|
||||
batch_np = tensor_to_numpy_uint8(batch)
|
||||
# Yield the whole batch at once to optimize pipe writes
|
||||
yield np.ascontiguousarray(batch_np)
|
||||
else:
|
||||
# Fallback for list input (e.g. pingpong or mixed sources)
|
||||
# We process individually as stacking might be expensive if they are not already contiguous tensors
|
||||
for img in image_sequence:
|
||||
yield np.ascontiguousarray(tensor_to_numpy_uint8(img))
|
||||
@@ -0,0 +1,529 @@
|
||||
"""
|
||||
Video encoding utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides FFmpeg-based video encoding with fallback to PIL for GIF/WebP.
|
||||
"""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from typing import List, Tuple, Optional, Iterator, Any, Callable
|
||||
from uuid import uuid4
|
||||
import numpy as np
|
||||
|
||||
|
||||
def detect_ffmpeg() -> Optional[str]:
|
||||
"""
|
||||
Detect FFmpeg binary location.
|
||||
|
||||
Returns:
|
||||
Path to FFmpeg executable, or None if not found
|
||||
"""
|
||||
ffmpeg_path = None
|
||||
|
||||
# Try imageio-ffmpeg first (common in Python environments)
|
||||
try:
|
||||
import imageio_ffmpeg
|
||||
ffmpeg_path = imageio_ffmpeg.get_ffmpeg_exe()
|
||||
print(f"Found ffmpeg via imageio_ffmpeg: {ffmpeg_path}")
|
||||
return ffmpeg_path
|
||||
except (ImportError, Exception):
|
||||
pass
|
||||
|
||||
# Try system PATH
|
||||
try:
|
||||
from shutil import which
|
||||
ffmpeg_path = which("ffmpeg")
|
||||
if ffmpeg_path:
|
||||
print(f"Found ffmpeg in system path: {ffmpeg_path}")
|
||||
return ffmpeg_path
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class FFmpegEncoder:
|
||||
"""
|
||||
FFmpeg-based video encoder supporting multiple formats.
|
||||
"""
|
||||
|
||||
def __init__(self, ffmpeg_path: Optional[str] = None):
|
||||
"""
|
||||
Initialize the encoder.
|
||||
|
||||
Args:
|
||||
ffmpeg_path: Path to FFmpeg executable (auto-detected if None)
|
||||
"""
|
||||
self.ffmpeg_path = ffmpeg_path or detect_ffmpeg()
|
||||
if not self.ffmpeg_path:
|
||||
raise RuntimeError("FFmpeg not found. Install ffmpeg or imageio-ffmpeg.")
|
||||
|
||||
def encode(
|
||||
self,
|
||||
images: List[np.ndarray],
|
||||
output_path: str,
|
||||
format_ext: str,
|
||||
frame_rate: float = 24.0,
|
||||
quality: int = 85,
|
||||
lossless: bool = False,
|
||||
loop_count: int = 0,
|
||||
codec: Optional[str] = None,
|
||||
progress_callback: Optional[Callable[[int, int], None]] = None
|
||||
) -> str:
|
||||
"""
|
||||
Encode images to video using FFmpeg.
|
||||
|
||||
Args:
|
||||
images: List of numpy arrays (H, W, C) in uint8 format
|
||||
output_path: Output file path
|
||||
format_ext: Output format extension (mp4, webm, gif)
|
||||
frame_rate: Frame rate in FPS
|
||||
quality: Quality level 1-100
|
||||
lossless: Use lossless encoding if supported
|
||||
loop_count: Number of loops (0 = infinite for GIF)
|
||||
codec: Specific codec to use (h264, h265, vp9, etc.)
|
||||
progress_callback: Optional callback(current, total) for progress
|
||||
|
||||
Returns:
|
||||
Path to the encoded file
|
||||
"""
|
||||
if not images:
|
||||
raise ValueError("No images provided for encoding")
|
||||
|
||||
# Get dimensions from first image
|
||||
height, width = images[0].shape[:2]
|
||||
has_alpha = images[0].shape[2] == 4 if len(images[0].shape) > 2 else False
|
||||
|
||||
# Determine input pixel format
|
||||
i_pix_fmt = "rgba" if has_alpha else "rgb24"
|
||||
dimensions = f"{width}x{height}"
|
||||
|
||||
# Build FFmpeg arguments
|
||||
args = self._build_ffmpeg_args(
|
||||
format_ext=format_ext,
|
||||
dimensions=dimensions,
|
||||
frame_rate=frame_rate,
|
||||
quality=quality,
|
||||
lossless=lossless,
|
||||
loop_count=loop_count,
|
||||
i_pix_fmt=i_pix_fmt,
|
||||
has_alpha=has_alpha,
|
||||
codec=codec,
|
||||
output_path=output_path
|
||||
)
|
||||
|
||||
# Execute encoding
|
||||
self._execute_encoding(args, images, i_pix_fmt, progress_callback)
|
||||
|
||||
return output_path
|
||||
|
||||
def _build_ffmpeg_args(
|
||||
self,
|
||||
format_ext: str,
|
||||
dimensions: str,
|
||||
frame_rate: float,
|
||||
quality: int,
|
||||
lossless: bool,
|
||||
loop_count: int,
|
||||
i_pix_fmt: str,
|
||||
has_alpha: bool,
|
||||
codec: Optional[str],
|
||||
output_path: str
|
||||
) -> List[str]:
|
||||
"""Build FFmpeg command arguments."""
|
||||
|
||||
# Loop arguments
|
||||
loop_args = []
|
||||
if format_ext == "gif":
|
||||
loop_args = ["-loop", "0" if loop_count == 0 else str(loop_count)]
|
||||
|
||||
# Base input arguments
|
||||
args = [
|
||||
self.ffmpeg_path, "-v", "error",
|
||||
"-f", "rawvideo",
|
||||
"-pix_fmt", i_pix_fmt,
|
||||
"-s", dimensions,
|
||||
"-r", str(frame_rate),
|
||||
"-i", "-"
|
||||
] + loop_args
|
||||
|
||||
# Format-specific encoding arguments
|
||||
if format_ext == "gif":
|
||||
args.extend(self._get_gif_args(quality))
|
||||
elif format_ext == "mp4":
|
||||
args.extend(self._get_mp4_args(quality, lossless, codec))
|
||||
elif format_ext == "webm":
|
||||
args.extend(self._get_webm_args(quality, lossless, has_alpha))
|
||||
else:
|
||||
# Default to MP4-like encoding
|
||||
args.extend(self._get_mp4_args(quality, lossless, codec))
|
||||
|
||||
# Add output path
|
||||
args.extend(["-y", output_path])
|
||||
|
||||
return args
|
||||
|
||||
def _get_gif_args(self, quality: int) -> List[str]:
|
||||
"""Get FFmpeg arguments for GIF encoding."""
|
||||
# Use palettegen for better quality
|
||||
if quality >= 80:
|
||||
return [
|
||||
"-vf", "split[s0][s1];[s0]palettegen=max_colors=256:stats_mode=diff[p];[s1][p]paletteuse=dither=sierra2",
|
||||
"-f", "gif"
|
||||
]
|
||||
else:
|
||||
return [
|
||||
"-vf", "split[s0][s1];[s0]palettegen[p];[s1][p]paletteuse",
|
||||
"-f", "gif"
|
||||
]
|
||||
|
||||
def _get_mp4_args(self, quality: int, lossless: bool, codec: Optional[str]) -> List[str]:
|
||||
"""Get FFmpeg arguments for MP4 encoding."""
|
||||
args = []
|
||||
|
||||
# Determine codec
|
||||
use_h265 = codec == "h265" or codec == "hevc"
|
||||
|
||||
if lossless:
|
||||
if use_h265:
|
||||
args.extend(["-c:v", "libx265", "-x265-params", "lossless=1"])
|
||||
else:
|
||||
args.extend(["-c:v", "libx264", "-crf", "0"])
|
||||
else:
|
||||
# Map quality (1-100) to CRF (51-0 for h264, lower is better)
|
||||
crf = int(51 - (quality / 100 * 33)) # Maps 1->51, 100->18
|
||||
|
||||
if use_h265:
|
||||
args.extend(["-c:v", "libx265", "-crf", str(crf + 5)]) # H.265 uses different CRF scale
|
||||
else:
|
||||
args.extend(["-c:v", "libx264", "-crf", str(crf)])
|
||||
|
||||
# Always use yuv420p for Discord compatibility
|
||||
args.extend(["-pix_fmt", "yuv420p", "-movflags", "faststart"])
|
||||
|
||||
return args
|
||||
|
||||
def _get_webm_args(self, quality: int, lossless: bool, has_alpha: bool) -> List[str]:
|
||||
"""Get FFmpeg arguments for WebM encoding."""
|
||||
args = ["-c:v", "libvpx-vp9"]
|
||||
|
||||
if lossless:
|
||||
args.extend(["-lossless", "1"])
|
||||
else:
|
||||
# Map quality to CRF (63-0 for VP9)
|
||||
crf = int(63 - (quality / 100 * 33)) # Maps 1->63, 100->30
|
||||
args.extend(["-crf", str(crf), "-b:v", "0"])
|
||||
|
||||
# Pixel format - support alpha if present
|
||||
pix_fmt = "yuva420p" if has_alpha else "yuv420p"
|
||||
args.extend(["-pix_fmt", pix_fmt])
|
||||
|
||||
# VP9 threading
|
||||
args.extend(["-row-mt", "1"])
|
||||
|
||||
return args
|
||||
|
||||
def _execute_encoding(
|
||||
self,
|
||||
args: List[str],
|
||||
images: List[np.ndarray],
|
||||
i_pix_fmt: str,
|
||||
progress_callback: Optional[Callable[[int, int], None]] = None
|
||||
) -> None:
|
||||
"""Execute FFmpeg process and feed frames."""
|
||||
total_frames = len(images)
|
||||
|
||||
# Create image chunk iterator for memory efficiency
|
||||
def image_chunks() -> Iterator[bytes]:
|
||||
for i, img in enumerate(images):
|
||||
# Ensure contiguous array for subprocess
|
||||
chunk = np.ascontiguousarray(img)
|
||||
if progress_callback:
|
||||
progress_callback(i + 1, total_frames)
|
||||
yield chunk.tobytes()
|
||||
|
||||
# Start FFmpeg process
|
||||
process = subprocess.Popen(
|
||||
args,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE
|
||||
)
|
||||
|
||||
# Feed frames
|
||||
try:
|
||||
for chunk in image_chunks():
|
||||
process.stdin.write(chunk)
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
|
||||
if process.returncode != 0:
|
||||
stderr = process.stderr.read().decode('utf-8', errors='ignore')
|
||||
raise RuntimeError(f"FFmpeg encoding failed: {stderr}")
|
||||
finally:
|
||||
if process.stdin:
|
||||
process.stdin.close()
|
||||
if process.stdout:
|
||||
process.stdout.close()
|
||||
if process.stderr:
|
||||
process.stderr.close()
|
||||
|
||||
|
||||
class PILEncoder:
|
||||
"""
|
||||
PIL-based encoder for GIF and WebP formats.
|
||||
Fallback when FFmpeg is not available.
|
||||
"""
|
||||
|
||||
def encode(
|
||||
self,
|
||||
images: List[Any], # PIL Images or numpy arrays
|
||||
output_path: str,
|
||||
format_ext: str,
|
||||
frame_rate: float = 24.0,
|
||||
quality: int = 85,
|
||||
lossless: bool = False,
|
||||
loop_count: int = 0,
|
||||
tensor_to_numpy_func: Optional[Callable] = None
|
||||
) -> str:
|
||||
"""
|
||||
Encode images using PIL.
|
||||
|
||||
Args:
|
||||
images: List of PIL Images or numpy arrays
|
||||
output_path: Output file path
|
||||
format_ext: Output format (gif, webp)
|
||||
frame_rate: Frame rate in FPS
|
||||
quality: Quality level 1-100
|
||||
lossless: Use lossless encoding for WebP
|
||||
loop_count: Number of loops (0 = infinite)
|
||||
tensor_to_numpy_func: Optional function to convert tensors to numpy
|
||||
|
||||
Returns:
|
||||
Path to the encoded file
|
||||
"""
|
||||
from PIL import Image
|
||||
|
||||
# Convert to PIL images if needed
|
||||
pil_images = []
|
||||
for img in images:
|
||||
if hasattr(img, 'shape'): # numpy array or tensor
|
||||
if tensor_to_numpy_func and hasattr(img, 'cpu'):
|
||||
img = tensor_to_numpy_func(img)
|
||||
elif hasattr(img, 'numpy'):
|
||||
img = img.numpy()
|
||||
pil_images.append(Image.fromarray(img.astype(np.uint8)))
|
||||
else:
|
||||
pil_images.append(img)
|
||||
|
||||
if not pil_images:
|
||||
raise ValueError("No images provided for encoding")
|
||||
|
||||
# Calculate frame duration in milliseconds
|
||||
duration = int(1000 / frame_rate)
|
||||
|
||||
if format_ext.lower() == "gif":
|
||||
self._encode_gif(pil_images, output_path, duration, loop_count)
|
||||
elif format_ext.lower() == "webp":
|
||||
self._encode_webp(pil_images, output_path, duration, loop_count, quality, lossless)
|
||||
else:
|
||||
# Single frame fallback
|
||||
pil_images[0].save(output_path, format=format_ext.upper())
|
||||
|
||||
return output_path
|
||||
|
||||
def _encode_gif(
|
||||
self,
|
||||
images: List[Any],
|
||||
output_path: str,
|
||||
duration: int,
|
||||
loop_count: int
|
||||
) -> None:
|
||||
"""Encode images as GIF."""
|
||||
durations = [duration] * len(images)
|
||||
images[0].save(
|
||||
output_path,
|
||||
format="GIF",
|
||||
append_images=images[1:] if len(images) > 1 else [],
|
||||
save_all=True,
|
||||
duration=durations,
|
||||
loop=0 if loop_count == 0 else loop_count,
|
||||
optimize=False
|
||||
)
|
||||
|
||||
def _encode_webp(
|
||||
self,
|
||||
images: List[Any],
|
||||
output_path: str,
|
||||
duration: int,
|
||||
loop_count: int,
|
||||
quality: int,
|
||||
lossless: bool
|
||||
) -> None:
|
||||
"""Encode images as WebP."""
|
||||
save_kwargs = {
|
||||
"format": "WEBP",
|
||||
"append_images": images[1:] if len(images) > 1 else [],
|
||||
"save_all": True,
|
||||
"duration": duration,
|
||||
"loop": 0 if loop_count == 0 else loop_count,
|
||||
}
|
||||
|
||||
if lossless:
|
||||
save_kwargs["lossless"] = True
|
||||
else:
|
||||
save_kwargs["quality"] = quality
|
||||
|
||||
images[0].save(output_path, **save_kwargs)
|
||||
|
||||
|
||||
def optimize_video_for_discord(
|
||||
input_file: str,
|
||||
ffmpeg_path: str,
|
||||
temp_dir: str
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Create a Discord-optimized version of a video file.
|
||||
|
||||
Args:
|
||||
input_file: Path to the input video file
|
||||
ffmpeg_path: Path to FFmpeg executable
|
||||
temp_dir: Directory for temporary files
|
||||
|
||||
Returns:
|
||||
Path to the optimized file, or None if optimization failed
|
||||
"""
|
||||
format_ext = os.path.splitext(input_file)[1].lstrip('.').lower()
|
||||
|
||||
# Use mkstemp for secure temporary file creation with restricted permissions (0600)
|
||||
# This prevents race conditions and ensures other users can't read the temp file
|
||||
fd, discord_optimized_file = tempfile.mkstemp(
|
||||
suffix=f".{format_ext}",
|
||||
prefix="discord_optimized_",
|
||||
dir=temp_dir
|
||||
)
|
||||
os.close(fd) # Close file descriptor immediately so FFmpeg can write to it
|
||||
|
||||
success = False
|
||||
try:
|
||||
if format_ext == "mp4":
|
||||
optimize_args = [
|
||||
ffmpeg_path, "-i", input_file,
|
||||
"-c:v", "libx264", "-pix_fmt", "yuv420p",
|
||||
"-movflags", "faststart", "-preset", "fast",
|
||||
"-profile:v", "baseline", "-level", "3.0",
|
||||
"-crf", "23",
|
||||
"-c:a", "aac", "-b:a", "128k",
|
||||
"-y", discord_optimized_file
|
||||
]
|
||||
elif format_ext == "webm":
|
||||
optimize_args = [
|
||||
ffmpeg_path, "-i", input_file,
|
||||
"-c:v", "libvpx-vp9",
|
||||
"-pix_fmt", "yuv420p",
|
||||
"-crf", "30", "-b:v", "0",
|
||||
"-deadline", "good",
|
||||
"-c:a", "libopus", "-b:a", "96k",
|
||||
"-y", discord_optimized_file
|
||||
]
|
||||
elif format_ext == "gif":
|
||||
optimize_args = [
|
||||
ffmpeg_path, "-i", input_file,
|
||||
"-vf", "fps=15,scale=trunc(iw/2)*2:trunc(ih/2)*2",
|
||||
"-y", discord_optimized_file
|
||||
]
|
||||
else:
|
||||
print(f"No optimization rules for format: {format_ext}")
|
||||
return None
|
||||
|
||||
print(f"Creating Discord-optimized version of {format_ext.upper()} file...")
|
||||
result = subprocess.run(
|
||||
optimize_args,
|
||||
capture_output=True,
|
||||
text=True
|
||||
)
|
||||
|
||||
if result.returncode == 0 and os.path.exists(discord_optimized_file):
|
||||
print(f"Discord-optimized file created: {discord_optimized_file}")
|
||||
success = True
|
||||
return discord_optimized_file
|
||||
else:
|
||||
print(f"Optimization failed: {result.stderr}")
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during Discord optimization: {e}")
|
||||
return None
|
||||
|
||||
finally:
|
||||
# Clean up temp file if optimization failed or wasn't supported
|
||||
if not success and os.path.exists(discord_optimized_file):
|
||||
try:
|
||||
os.remove(discord_optimized_file)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def mux_audio_to_video(
|
||||
video_path: str,
|
||||
audio_waveform: np.ndarray,
|
||||
sample_rate: int,
|
||||
format_ext: str,
|
||||
ffmpeg_path: str,
|
||||
output_path: str,
|
||||
channels: int = 2
|
||||
) -> bool:
|
||||
"""
|
||||
Mux audio into a video file.
|
||||
|
||||
Args:
|
||||
video_path: Path to the video file
|
||||
audio_waveform: Audio data as numpy array
|
||||
sample_rate: Audio sample rate
|
||||
format_ext: Video format extension
|
||||
ffmpeg_path: Path to FFmpeg executable
|
||||
output_path: Output path for the muxed file
|
||||
channels: Number of audio channels
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
"""
|
||||
try:
|
||||
# Determine audio codec based on format
|
||||
if format_ext == "mp4":
|
||||
audio_pass = ["-c:a", "aac", "-b:a", "192k"]
|
||||
elif format_ext == "webm":
|
||||
audio_pass = ["-c:a", "libopus", "-b:a", "128k"]
|
||||
else:
|
||||
audio_pass = ["-c:a", "libopus", "-b:a", "128k"]
|
||||
|
||||
mux_args = [
|
||||
ffmpeg_path, "-v", "error", "-y",
|
||||
"-i", video_path,
|
||||
"-ar", str(sample_rate),
|
||||
"-ac", str(channels),
|
||||
"-f", "f32le",
|
||||
"-i", "-",
|
||||
"-c:v", "copy"
|
||||
] + audio_pass + ["-shortest", output_path]
|
||||
|
||||
# Ensure contiguous array for subprocess
|
||||
audio_data = np.ascontiguousarray(audio_waveform)
|
||||
|
||||
result = subprocess.run(
|
||||
mux_args,
|
||||
input=memoryview(audio_data),
|
||||
capture_output=True
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
print(f"Successfully muxed audio to video: {output_path}")
|
||||
return True
|
||||
else:
|
||||
print(f"Audio muxing failed: {result.stderr.decode('utf-8', errors='ignore')}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error muxing audio: {e}")
|
||||
return False
|
||||
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
Path utilities for ComfyUI-DiscordSend
|
||||
|
||||
Provides functions for handling output directories and file paths.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def get_output_directory(
|
||||
save_output: bool,
|
||||
comfy_output_dir: str,
|
||||
temp_dir: str,
|
||||
subfolder: str = "discord_output"
|
||||
) -> str:
|
||||
"""
|
||||
Determine the appropriate output directory based on save settings.
|
||||
|
||||
Args:
|
||||
save_output: Whether files should be saved permanently
|
||||
comfy_output_dir: ComfyUI's output directory path
|
||||
temp_dir: ComfyUI's temporary directory path
|
||||
subfolder: Subfolder name within output directory (default: "discord_output")
|
||||
|
||||
Returns:
|
||||
Path to the destination directory
|
||||
"""
|
||||
if save_output:
|
||||
# Create output subfolder in the ComfyUI output directory
|
||||
dest_folder = os.path.join(comfy_output_dir, subfolder)
|
||||
os.makedirs(dest_folder, exist_ok=True)
|
||||
else:
|
||||
# Use ComfyUI's temporary directory for preview-only files
|
||||
dest_folder = temp_dir
|
||||
os.makedirs(dest_folder, exist_ok=True)
|
||||
print(f"Using temporary directory for preview: {dest_folder}")
|
||||
|
||||
return dest_folder
|
||||
|
||||
|
||||
def ensure_directory_exists(path: str) -> str:
|
||||
"""
|
||||
Ensure a directory exists, creating it if necessary.
|
||||
|
||||
Args:
|
||||
path: Directory path to ensure exists
|
||||
|
||||
Returns:
|
||||
The same path (for chaining)
|
||||
"""
|
||||
os.makedirs(path, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def get_unique_filepath(
|
||||
directory: str,
|
||||
filename: str,
|
||||
extension: str,
|
||||
counter: Optional[int] = None
|
||||
) -> str:
|
||||
"""
|
||||
Generate a unique filepath, optionally with a counter.
|
||||
|
||||
Args:
|
||||
directory: Base directory
|
||||
filename: Base filename (without extension)
|
||||
extension: File extension (with or without leading dot)
|
||||
counter: Optional counter to append to filename
|
||||
|
||||
Returns:
|
||||
Full filepath
|
||||
"""
|
||||
# Ensure extension has leading dot
|
||||
if not extension.startswith("."):
|
||||
extension = "." + extension
|
||||
|
||||
if counter is not None:
|
||||
full_filename = f"{filename}_{counter:05d}{extension}"
|
||||
else:
|
||||
full_filename = f"{filename}{extension}"
|
||||
|
||||
return os.path.join(directory, full_filename)
|
||||
|
||||
|
||||
def validate_path_is_safe(path: str, base_dir: Optional[str] = None) -> None:
|
||||
"""
|
||||
Validate that a path is safe to write to.
|
||||
|
||||
Checks:
|
||||
- If base_dir is provided, path is contained within base_dir (to prevent ../ traversal)
|
||||
- Path is not a symlink (to prevent overwriting targets)
|
||||
- Parent directories are not symlinks (to prevent path traversal via symlinks)
|
||||
|
||||
Args:
|
||||
path: File path to validate
|
||||
base_dir: Optional base directory to restrict path to
|
||||
|
||||
Raises:
|
||||
ValueError: If path is unsafe
|
||||
"""
|
||||
# Check if path is within base_dir
|
||||
if base_dir:
|
||||
abs_base = os.path.abspath(base_dir)
|
||||
abs_path = os.path.abspath(path)
|
||||
|
||||
# Use commonpath to ensure path is within base_dir
|
||||
# We need to handle potential different drives on Windows which raises ValueError
|
||||
try:
|
||||
common = os.path.commonpath([abs_base, abs_path])
|
||||
except ValueError:
|
||||
# Raised if paths are on different drives
|
||||
raise ValueError(f"Security error: Path '{path}' is on a different drive than allowed directory '{base_dir}'.")
|
||||
|
||||
if common != abs_base:
|
||||
raise ValueError(f"Security error: Path '{path}' is outside the allowed directory '{base_dir}'.")
|
||||
|
||||
# Check if path exists and is a symlink
|
||||
if os.path.islink(path):
|
||||
raise ValueError(f"Security error: Output path '{path}' is a symlink. Overwriting symlinks is not allowed.")
|
||||
|
||||
# Verify parent directories
|
||||
# Walk up the tree to find the first existing directory
|
||||
current_dir = os.path.dirname(os.path.abspath(path))
|
||||
|
||||
# Safety valve to prevent infinite loops (though OS paths are finite)
|
||||
# We check existence. If it doesn't exist, we check if it's a broken symlink (islink returns True even for broken links)
|
||||
# Then move to parent.
|
||||
|
||||
while current_dir and current_dir != os.path.dirname(current_dir): # Until root
|
||||
if os.path.islink(current_dir):
|
||||
raise ValueError(f"Security error: Path component '{current_dir}' is a symlink. Writing through directory symlinks is not allowed.")
|
||||
|
||||
if os.path.exists(current_dir):
|
||||
# Once we hit an existing directory, we verify it matches its realpath
|
||||
# This catches hidden symlinks further up that might have been resolved by abspath but diverge in realpath
|
||||
real_dir = os.path.realpath(current_dir)
|
||||
abs_dir = os.path.abspath(current_dir)
|
||||
|
||||
if real_dir != abs_dir:
|
||||
raise ValueError(f"Security error: Path resolution mismatch for '{current_dir}'. "
|
||||
f"Symlinks in output paths are not allowed (Real: {real_dir}, Abs: {abs_dir}).")
|
||||
# If the existing ancestor is safe, we assume children created under it will be normal directories
|
||||
# (unless we have a race condition, but we can't solve that fully without openat)
|
||||
break
|
||||
|
||||
current_dir = os.path.dirname(current_dir)
|
||||
@@ -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',
|
||||
]
|
||||
@@ -15,17 +15,25 @@ NEGATIVE_INDICATORS = [
|
||||
"extra limbs", "bad anatomy", "watermark", "text", "signature"
|
||||
]
|
||||
|
||||
# Node types that can contain prompts
|
||||
PROMPT_NODE_TYPES = [
|
||||
"CLIPTextEncode", # Standard SD 1.5 prompt node
|
||||
"SDXLPromptEncoder", # SDXL prompt encoder
|
||||
"SDXLTextEncode", # Another SDXL text node
|
||||
]
|
||||
|
||||
|
||||
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
|
||||
|
||||
Analyzes ComfyUI workflow structure to find prompt nodes (CLIPTextEncode,
|
||||
SDXLPromptEncoder, SDXLTextEncode, and other text encoding 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
|
||||
"""
|
||||
@@ -47,25 +55,25 @@ def extract_prompts_from_workflow(workflow_data: Any) -> Tuple[Optional[str], Op
|
||||
positive_prompt = None
|
||||
negative_prompt = None
|
||||
|
||||
# Find CLIPTextEncode nodes
|
||||
# Find prompt nodes (CLIPTextEncode, SDXL nodes, etc.)
|
||||
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:
|
||||
|
||||
prompt_nodes = _find_prompt_nodes(nodes)
|
||||
|
||||
if not prompt_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)
|
||||
if len(prompt_nodes) == 1:
|
||||
# Single prompt node - assume it's the positive prompt
|
||||
positive_prompt = _get_prompt_text(prompt_nodes[0])
|
||||
elif len(prompt_nodes) >= 2:
|
||||
# Multiple prompt nodes - need to determine which is which
|
||||
positive_prompt, negative_prompt = _classify_prompts(prompt_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:
|
||||
@@ -74,41 +82,47 @@ def extract_prompts_from_workflow(workflow_data: Any) -> Tuple[Optional[str], Op
|
||||
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 = []
|
||||
|
||||
def _find_prompt_nodes(nodes: Union[List, Dict]) -> List[Dict]:
|
||||
"""Find all prompt nodes (CLIPTextEncode, SDXL nodes, etc.) in the workflow."""
|
||||
prompt_nodes = []
|
||||
|
||||
if isinstance(nodes, list):
|
||||
for node in nodes:
|
||||
if _is_clip_text_encode(node):
|
||||
clip_nodes.append(node)
|
||||
if _is_prompt_node(node):
|
||||
prompt_nodes.append(node)
|
||||
elif isinstance(nodes, dict):
|
||||
for node_id, node in nodes.items():
|
||||
if _is_clip_text_encode(node):
|
||||
if _is_prompt_node(node):
|
||||
node_copy = dict(node)
|
||||
node_copy["id"] = node_id
|
||||
clip_nodes.append(node_copy)
|
||||
|
||||
return clip_nodes
|
||||
prompt_nodes.append(node_copy)
|
||||
|
||||
return prompt_nodes
|
||||
|
||||
|
||||
def _is_clip_text_encode(node: Any) -> bool:
|
||||
"""Check if a node is a CLIPTextEncode node with valid text."""
|
||||
def _is_prompt_node(node: Any) -> bool:
|
||||
"""Check if a node is a prompt node (CLIPTextEncode, SDXL, etc.) 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 against known prompt node types
|
||||
if node_type not in PROMPT_NODE_TYPES:
|
||||
# Also check for dynamic text/prompt nodes (e.g., custom nodes)
|
||||
if node_type and ("Text" in node_type and ("Encode" in node_type or "Prompt" in node_type)):
|
||||
pass # Allow these through
|
||||
else:
|
||||
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."""
|
||||
"""Extract the prompt text from a prompt node."""
|
||||
# Workflow format (widgets_values)
|
||||
widgets = node.get("widgets_values", [])
|
||||
if isinstance(widgets, list) and len(widgets) > 0 and isinstance(widgets[0], str):
|
||||
@@ -122,20 +136,20 @@ def _get_prompt_text(node: Dict) -> Optional[str]:
|
||||
return None
|
||||
|
||||
|
||||
def _classify_prompts(clip_nodes: List[Dict], workflow_data: Dict) -> Tuple[Optional[str], Optional[str]]:
|
||||
def _classify_prompts(prompt_nodes: List[Dict], workflow_data: Dict) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Classify which CLIP nodes contain positive vs negative prompts.
|
||||
|
||||
Classify which prompt 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:
|
||||
if not prompt_nodes:
|
||||
return None, None
|
||||
|
||||
# First pass: Score all nodes based on content
|
||||
node_scores = []
|
||||
for node in clip_nodes:
|
||||
for node in prompt_nodes:
|
||||
prompt_text = _get_prompt_text(node)
|
||||
# Skip empty or None text
|
||||
if not prompt_text or not prompt_text.strip():
|
||||
@@ -169,14 +183,14 @@ def _classify_prompts(clip_nodes: List[Dict], workflow_data: Dict) -> Tuple[Opti
|
||||
else:
|
||||
# All scores are 0, use connection analysis
|
||||
positive_prompt, negative_prompt = _classify_by_connections(
|
||||
clip_nodes, workflow_data, None, None
|
||||
prompt_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:
|
||||
if positive_prompt is None and negative_prompt is None and len(prompt_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])
|
||||
positive_prompt = _get_prompt_text(prompt_nodes[0])
|
||||
negative_prompt = _get_prompt_text(prompt_nodes[1])
|
||||
elif positive_prompt is None and negative_prompt is not None:
|
||||
# Find the other prompt
|
||||
for _, _, text in node_scores:
|
||||
@@ -194,7 +208,7 @@ def _classify_prompts(clip_nodes: List[Dict], workflow_data: Dict) -> Tuple[Opti
|
||||
|
||||
|
||||
def _classify_by_connections(
|
||||
clip_nodes: List[Dict],
|
||||
prompt_nodes: List[Dict],
|
||||
workflow_data: Dict,
|
||||
existing_positive: Optional[str],
|
||||
existing_negative: Optional[str]
|
||||
@@ -235,7 +249,7 @@ def _classify_by_connections(
|
||||
to_slot = link[3]
|
||||
|
||||
# Find matching CLIP node and sampler
|
||||
for clip_node in clip_nodes:
|
||||
for clip_node in prompt_nodes:
|
||||
clip_id = clip_node.get("id")
|
||||
if clip_id == from_node_id:
|
||||
for sampler in samplers:
|
||||
@@ -11,18 +11,19 @@ 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",
|
||||
]
|
||||
# Pre-compile regex for faster matching
|
||||
WEBHOOK_REGEX = re.compile(
|
||||
r"(discord\.com/api/webhooks|discordapp\.com/api/webhooks)",
|
||||
re.IGNORECASE
|
||||
)
|
||||
|
||||
GITHUB_TOKEN_PREFIXES = [
|
||||
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:
|
||||
@@ -30,9 +31,8 @@ def is_webhook_url(value: str) -> bool:
|
||||
if not isinstance(value, str):
|
||||
return False
|
||||
|
||||
for pattern in WEBHOOK_PATTERNS:
|
||||
if re.search(pattern, value, re.IGNORECASE):
|
||||
return True
|
||||
if WEBHOOK_REGEX.search(value):
|
||||
return True
|
||||
|
||||
# Also check for generic webhook patterns in URLs
|
||||
if value.startswith("http") and ("webhook" in value.lower() or "discord" in value.lower()):
|
||||
@@ -46,11 +46,7 @@ def is_github_token(value: str) -> bool:
|
||||
if not isinstance(value, str):
|
||||
return False
|
||||
|
||||
for prefix in GITHUB_TOKEN_PREFIXES:
|
||||
if value.startswith(prefix):
|
||||
return True
|
||||
|
||||
return False
|
||||
return value.startswith(GITHUB_TOKEN_PREFIXES)
|
||||
|
||||
|
||||
def is_potential_token(value: str, context_type: str = "") -> bool:
|
||||
@@ -139,29 +135,45 @@ def sanitize_node_inputs(inputs: Dict, node_type: str = "") -> Dict:
|
||||
return result
|
||||
|
||||
|
||||
def sanitize_node(node: Dict) -> Dict:
|
||||
def sanitize_node(node: Any) -> Any:
|
||||
"""
|
||||
Sanitize a single ComfyUI node.
|
||||
|
||||
Args:
|
||||
node: The node dictionary
|
||||
node: The node dictionary or value
|
||||
|
||||
Returns:
|
||||
Sanitized node dictionary
|
||||
Sanitized node dictionary or value
|
||||
"""
|
||||
if not isinstance(node, dict):
|
||||
if isinstance(node, str):
|
||||
return sanitize_string(node)
|
||||
return node
|
||||
|
||||
result = dict(node)
|
||||
node_type = result.get("type", "")
|
||||
result = {}
|
||||
node_type = node.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)
|
||||
for key, value in node.items():
|
||||
# Handle known sensitive keys
|
||||
if key in ("webhook_url", "github_token"):
|
||||
result[key] = ""
|
||||
continue
|
||||
|
||||
# Context-aware sanitization for inputs and widgets
|
||||
if key == "inputs" and isinstance(value, dict):
|
||||
result[key] = sanitize_node_inputs(value, node_type)
|
||||
elif key == "widgets_values" and isinstance(value, list):
|
||||
result[key] = sanitize_widget_values(value, node_type)
|
||||
else:
|
||||
# Generic sanitization for other fields
|
||||
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
|
||||
|
||||
return result
|
||||
|
||||
@@ -178,12 +190,28 @@ def sanitize_dict(data: Dict) -> Dict:
|
||||
"""
|
||||
result = {}
|
||||
|
||||
# Check if this is a workflow object with nodes
|
||||
# We want to process nodes specifically using sanitize_node to ensure correct context
|
||||
# and avoid double-processing (once as generic dict/list, once as nodes)
|
||||
is_workflow = "nodes" in data and isinstance(data["nodes"], (list, dict))
|
||||
|
||||
for key, value in data.items():
|
||||
# Handle known sensitive keys
|
||||
if key in ("webhook_url", "github_token"):
|
||||
result[key] = ""
|
||||
continue
|
||||
|
||||
# Special handling for "nodes" in workflow
|
||||
if is_workflow and key == "nodes":
|
||||
if isinstance(value, list):
|
||||
result[key] = [sanitize_node(n) for n in value]
|
||||
elif isinstance(value, dict):
|
||||
result[key] = {k: sanitize_node(v) for k, v in value.items()}
|
||||
else:
|
||||
# Fallback if nodes is neither list nor dict (unlikely)
|
||||
result[key] = value
|
||||
continue
|
||||
|
||||
# Handle nested structures
|
||||
if isinstance(value, dict):
|
||||
result[key] = sanitize_dict(value)
|
||||
@@ -194,14 +222,6 @@ def sanitize_dict(data: Dict) -> Dict:
|
||||
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
|
||||
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""
|
||||
Pytest configuration file for tests.
|
||||
|
||||
This module handles test isolation by ensuring that real packages are imported
|
||||
before any mocking occurs, and by providing cleanup fixtures.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Add project root to path
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# Store references to real modules before any mocking
|
||||
# This ensures tests that need real numpy/PIL can use them
|
||||
_real_numpy = None
|
||||
_real_PIL = None
|
||||
|
||||
def pytest_configure(config):
|
||||
"""Called after command line options have been parsed and all plugins loaded."""
|
||||
global _real_numpy, _real_PIL
|
||||
|
||||
# Import real modules and store references
|
||||
try:
|
||||
import numpy
|
||||
_real_numpy = numpy
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
import PIL
|
||||
import PIL.Image
|
||||
import PIL.PngImagePlugin
|
||||
_real_PIL = PIL
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
def get_real_numpy():
|
||||
"""Get the real numpy module, not a mock."""
|
||||
if _real_numpy is None:
|
||||
import numpy
|
||||
return numpy
|
||||
return _real_numpy
|
||||
|
||||
|
||||
def get_real_PIL():
|
||||
"""Get the real PIL module, not a mock."""
|
||||
if _real_PIL is None:
|
||||
import PIL
|
||||
return PIL
|
||||
return _real_PIL
|
||||
@@ -3,10 +3,15 @@ import sys
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from discordsend_utils.discord_api import send_to_discord_with_retry
|
||||
from shared.discord import send_to_discord_with_retry
|
||||
|
||||
class TestDiscordAPI(unittest.TestCase):
|
||||
"""Tests for the Discord API utility with mocked network responses."""
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
"""Tests for shared/filename_utils.py"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from shared.filename_utils import build_filename_with_metadata, get_timestamp_string
|
||||
|
||||
|
||||
class TestBuildFilenameWithMetadata(unittest.TestCase):
|
||||
"""Test build_filename_with_metadata function."""
|
||||
|
||||
def test_prefix_only(self):
|
||||
"""Test with just a prefix, no metadata."""
|
||||
result, info = build_filename_with_metadata("image")
|
||||
self.assertEqual(result, "image")
|
||||
self.assertEqual(info, {})
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_with_date(self, mock_time):
|
||||
"""Test adding date to filename."""
|
||||
mock_time.strftime.return_value = "2026-01-20"
|
||||
result, info = build_filename_with_metadata("image", add_date=True)
|
||||
self.assertEqual(result, "image_2026-01-20")
|
||||
self.assertEqual(info["date"], "2026-01-20")
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_with_time(self, mock_time):
|
||||
"""Test adding time to filename."""
|
||||
mock_time.strftime.return_value = "14-30-00"
|
||||
result, info = build_filename_with_metadata("image", add_time=True)
|
||||
self.assertEqual(result, "image_14-30-00")
|
||||
self.assertEqual(info["time"], "14-30-00")
|
||||
|
||||
def test_with_dimensions(self):
|
||||
"""Test adding dimensions to filename."""
|
||||
result, info = build_filename_with_metadata(
|
||||
"image", add_dimensions=True, width=1920, height=1080
|
||||
)
|
||||
self.assertEqual(result, "image_1920x1080")
|
||||
self.assertEqual(info["dimensions"], "1920x1080")
|
||||
|
||||
def test_dimensions_without_values(self):
|
||||
"""Test that dimensions are not added without width/height."""
|
||||
result, info = build_filename_with_metadata("image", add_dimensions=True)
|
||||
self.assertEqual(result, "image")
|
||||
self.assertNotIn("dimensions", info)
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_all_metadata(self, mock_time):
|
||||
"""Test with all metadata options."""
|
||||
mock_time.strftime.side_effect = ["2026-01-20", "14-30-00"]
|
||||
result, info = build_filename_with_metadata(
|
||||
"output",
|
||||
add_date=True,
|
||||
add_time=True,
|
||||
add_dimensions=True,
|
||||
width=512,
|
||||
height=768,
|
||||
)
|
||||
self.assertEqual(result, "output_2026-01-20_14-30-00_512x768")
|
||||
self.assertEqual(info["date"], "2026-01-20")
|
||||
self.assertEqual(info["time"], "14-30-00")
|
||||
self.assertEqual(info["dimensions"], "512x768")
|
||||
|
||||
def test_with_existing_info_dict(self):
|
||||
"""Test that existing info_dict is updated, not replaced."""
|
||||
existing_info = {"existing_key": "existing_value"}
|
||||
result, info = build_filename_with_metadata(
|
||||
"image", add_dimensions=True, width=100, height=100, info_dict=existing_info
|
||||
)
|
||||
self.assertEqual(info["existing_key"], "existing_value")
|
||||
self.assertEqual(info["dimensions"], "100x100")
|
||||
self.assertIs(info, existing_info) # Same dict object
|
||||
|
||||
|
||||
class TestGetTimestampString(unittest.TestCase):
|
||||
"""Test get_timestamp_string function."""
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_date_only(self, mock_time):
|
||||
"""Test timestamp with date only."""
|
||||
mock_time.strftime.return_value = "2026-01-20"
|
||||
result = get_timestamp_string(include_date=True, include_time=False)
|
||||
self.assertEqual(result, "2026-01-20")
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_time_only(self, mock_time):
|
||||
"""Test timestamp with time only."""
|
||||
mock_time.strftime.return_value = "14-30-00"
|
||||
result = get_timestamp_string(include_date=False, include_time=True)
|
||||
self.assertEqual(result, "14-30-00")
|
||||
|
||||
@patch("shared.filename_utils.time")
|
||||
def test_both(self, mock_time):
|
||||
"""Test timestamp with both date and time."""
|
||||
mock_time.strftime.side_effect = ["2026-01-20", "14-30-00"]
|
||||
result = get_timestamp_string(include_date=True, include_time=True)
|
||||
self.assertEqual(result, "2026-01-20_14-30-00")
|
||||
|
||||
def test_neither(self):
|
||||
"""Test timestamp with neither date nor time."""
|
||||
result = get_timestamp_string(include_date=False, include_time=False)
|
||||
self.assertEqual(result, "")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Tests for shared/media/format_utils.py"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
import tempfile
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from shared.media.format_utils import (
|
||||
parse_format_string,
|
||||
normalize_video_extension,
|
||||
get_mime_type,
|
||||
validate_video_for_discord,
|
||||
is_animated_format,
|
||||
supports_alpha,
|
||||
)
|
||||
|
||||
|
||||
class TestParseFormatString(unittest.TestCase):
|
||||
"""Test parse_format_string function."""
|
||||
|
||||
def test_video_h264_mp4(self):
|
||||
"""Test parsing video/h264-mp4 format."""
|
||||
fmt_type, fmt_ext = parse_format_string("video/h264-mp4")
|
||||
self.assertEqual(fmt_type, "video")
|
||||
self.assertEqual(fmt_ext, "h264-mp4")
|
||||
|
||||
def test_image_gif(self):
|
||||
"""Test parsing image/gif format."""
|
||||
fmt_type, fmt_ext = parse_format_string("image/gif")
|
||||
self.assertEqual(fmt_type, "image")
|
||||
self.assertEqual(fmt_ext, "gif")
|
||||
|
||||
def test_simple_format(self):
|
||||
"""Test parsing simple format string without slash."""
|
||||
fmt_type, fmt_ext = parse_format_string("mp4")
|
||||
self.assertEqual(fmt_type, "video")
|
||||
self.assertEqual(fmt_ext, "mp4")
|
||||
|
||||
|
||||
class TestNormalizeVideoExtension(unittest.TestCase):
|
||||
"""Test normalize_video_extension function."""
|
||||
|
||||
def test_h264_mp4(self):
|
||||
"""Test normalizing h264-mp4 to mp4."""
|
||||
self.assertEqual(normalize_video_extension("video/h264-mp4"), "mp4")
|
||||
|
||||
def test_h265_mp4(self):
|
||||
"""Test normalizing h265-mp4 to mp4."""
|
||||
self.assertEqual(normalize_video_extension("video/h265-mp4"), "mp4")
|
||||
|
||||
def test_vp9_webm(self):
|
||||
"""Test normalizing vp9-webm to webm."""
|
||||
self.assertEqual(normalize_video_extension("video/vp9-webm"), "webm")
|
||||
|
||||
def test_prores(self):
|
||||
"""Test normalizing prores to mov."""
|
||||
self.assertEqual(normalize_video_extension("video/prores"), "mov")
|
||||
|
||||
def test_gif_passthrough(self):
|
||||
"""Test gif format passes through unchanged."""
|
||||
self.assertEqual(normalize_video_extension("image/gif"), "gif")
|
||||
|
||||
def test_unknown_passthrough(self):
|
||||
"""Test unknown format passes through unchanged."""
|
||||
self.assertEqual(normalize_video_extension("video/custom"), "custom")
|
||||
|
||||
|
||||
class TestGetMimeType(unittest.TestCase):
|
||||
"""Test get_mime_type function."""
|
||||
|
||||
def test_mp4(self):
|
||||
"""Test MIME type for mp4."""
|
||||
self.assertEqual(get_mime_type("mp4"), "video/mp4")
|
||||
|
||||
def test_webm(self):
|
||||
"""Test MIME type for webm."""
|
||||
self.assertEqual(get_mime_type("webm"), "video/webm")
|
||||
|
||||
def test_gif(self):
|
||||
"""Test MIME type for gif."""
|
||||
self.assertEqual(get_mime_type("gif"), "image/gif")
|
||||
|
||||
def test_mov(self):
|
||||
"""Test MIME type for mov."""
|
||||
self.assertEqual(get_mime_type("mov"), "video/quicktime")
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""Test MIME type lookup is case insensitive."""
|
||||
self.assertEqual(get_mime_type("MP4"), "video/mp4")
|
||||
|
||||
def test_unknown_format(self):
|
||||
"""Test unknown format returns octet-stream."""
|
||||
self.assertEqual(get_mime_type("xyz"), "application/octet-stream")
|
||||
|
||||
|
||||
class TestValidateVideoForDiscord(unittest.TestCase):
|
||||
"""Test validate_video_for_discord function."""
|
||||
|
||||
def test_nonexistent_file(self):
|
||||
"""Test validation of nonexistent file."""
|
||||
is_valid, msg = validate_video_for_discord("/nonexistent/file.mp4")
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("does not exist", msg)
|
||||
|
||||
def test_empty_file(self):
|
||||
"""Test validation of empty file."""
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as f:
|
||||
temp_path = f.name
|
||||
try:
|
||||
is_valid, msg = validate_video_for_discord(temp_path)
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("empty", msg)
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
def test_small_file(self):
|
||||
"""Test validation of suspiciously small file."""
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as f:
|
||||
f.write(b"x" * 100) # 100 bytes
|
||||
temp_path = f.name
|
||||
try:
|
||||
is_valid, msg = validate_video_for_discord(temp_path)
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("small", msg)
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
def test_valid_mp4(self):
|
||||
"""Test validation of valid mp4 file."""
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as f:
|
||||
f.write(b"x" * 10000) # 10KB
|
||||
temp_path = f.name
|
||||
try:
|
||||
is_valid, msg = validate_video_for_discord(temp_path)
|
||||
self.assertTrue(is_valid)
|
||||
self.assertEqual(msg, "Valid")
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
def test_valid_webm(self):
|
||||
"""Test validation of valid webm file."""
|
||||
with tempfile.NamedTemporaryFile(suffix=".webm", delete=False) as f:
|
||||
f.write(b"x" * 10000)
|
||||
temp_path = f.name
|
||||
try:
|
||||
is_valid, msg = validate_video_for_discord(temp_path)
|
||||
self.assertTrue(is_valid)
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
def test_mov_needs_conversion(self):
|
||||
"""Test that MOV files are flagged for conversion."""
|
||||
with tempfile.NamedTemporaryFile(suffix=".mov", delete=False) as f:
|
||||
f.write(b"x" * 10000)
|
||||
temp_path = f.name
|
||||
try:
|
||||
is_valid, msg = validate_video_for_discord(temp_path)
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("conversion", msg)
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
|
||||
class TestIsAnimatedFormat(unittest.TestCase):
|
||||
"""Test is_animated_format function."""
|
||||
|
||||
def test_animated_formats(self):
|
||||
"""Test formats that support animation."""
|
||||
animated = ["gif", "webp", "mp4", "webm", "mov", "avi", "mkv", "apng"]
|
||||
for fmt in animated:
|
||||
self.assertTrue(is_animated_format(fmt), f"{fmt} should be animated")
|
||||
|
||||
def test_static_formats(self):
|
||||
"""Test formats that don't support animation."""
|
||||
static = ["png", "jpg", "jpeg", "bmp"]
|
||||
for fmt in static:
|
||||
self.assertFalse(is_animated_format(fmt), f"{fmt} should not be animated")
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""Test case insensitivity."""
|
||||
self.assertTrue(is_animated_format("GIF"))
|
||||
self.assertTrue(is_animated_format("Mp4"))
|
||||
|
||||
|
||||
class TestSupportsAlpha(unittest.TestCase):
|
||||
"""Test supports_alpha function."""
|
||||
|
||||
def test_alpha_formats(self):
|
||||
"""Test formats that support alpha channel."""
|
||||
alpha = ["webm", "gif", "webp", "png", "apng", "mov"]
|
||||
for fmt in alpha:
|
||||
self.assertTrue(supports_alpha(fmt), f"{fmt} should support alpha")
|
||||
|
||||
def test_no_alpha_formats(self):
|
||||
"""Test formats that don't support alpha."""
|
||||
no_alpha = ["mp4", "jpg", "jpeg", "avi"]
|
||||
for fmt in no_alpha:
|
||||
self.assertFalse(supports_alpha(fmt), f"{fmt} should not support alpha")
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""Test case insensitivity."""
|
||||
self.assertTrue(supports_alpha("PNG"))
|
||||
self.assertTrue(supports_alpha("WebM"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,99 @@
|
||||
|
||||
import unittest
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# Mock torch and other heavy dependencies
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["PIL"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
sys.modules["folder_paths"] = MagicMock()
|
||||
sys.modules["comfy"] = MagicMock()
|
||||
sys.modules["comfy.cli_args"] = MagicMock()
|
||||
|
||||
# Now we can safely import
|
||||
from shared.github_integration import validate_github_repo, validate_file_path, update_github_cdn_urls
|
||||
|
||||
class TestGitHubValidation(unittest.TestCase):
|
||||
|
||||
def test_validate_github_repo_valid(self):
|
||||
"""Test valid GitHub repository formats."""
|
||||
valid_repos = [
|
||||
"username/repo",
|
||||
"user-name/repo-name",
|
||||
"user-name/repo.name", # Dot in repo is valid
|
||||
"user-name/repo_name", # Underscore in repo is valid
|
||||
"0123/4567"
|
||||
]
|
||||
for repo in valid_repos:
|
||||
with self.subTest(repo=repo):
|
||||
self.assertTrue(validate_github_repo(repo), f"Failed for {repo}")
|
||||
|
||||
def test_validate_github_repo_invalid(self):
|
||||
"""Test invalid GitHub repository formats (traversal, injection, invalid chars)."""
|
||||
invalid_repos = [
|
||||
"username/repo/../other", # Traversal
|
||||
"username/repo?query=1", # Query injection
|
||||
"username", # Missing slash
|
||||
"/repo", # Missing username
|
||||
"user/", # Missing repo
|
||||
"user/repo/", # Trailing slash (strict check)
|
||||
"../../user/repo", # Traversal at start
|
||||
"user/repo#fragment", # Fragment
|
||||
"user/repo;rm -rf", # Command injection style
|
||||
"user/repo\nnewline", # Newline
|
||||
"user.name/repo", # Dot in username (invalid)
|
||||
"user_name/repo", # Underscore in username (invalid)
|
||||
]
|
||||
for repo in invalid_repos:
|
||||
with self.subTest(repo=repo):
|
||||
self.assertFalse(validate_github_repo(repo), f"Should have failed for {repo}")
|
||||
|
||||
def test_validate_file_path_valid(self):
|
||||
"""Test valid file paths."""
|
||||
valid_paths = [
|
||||
"file.txt",
|
||||
"path/to/file.txt",
|
||||
"folder/subfolder/file.md",
|
||||
"README.md",
|
||||
"docs/image.png"
|
||||
]
|
||||
for path in valid_paths:
|
||||
with self.subTest(path=path):
|
||||
self.assertTrue(validate_file_path(path), f"Failed for {path}")
|
||||
|
||||
def test_validate_file_path_invalid(self):
|
||||
"""Test invalid file paths (traversal, absolute)."""
|
||||
invalid_paths = [
|
||||
"../file.txt", # Traversal
|
||||
"path/../file.txt", # Traversal inside
|
||||
"/etc/passwd", # Absolute path
|
||||
"/file.txt", # Absolute path
|
||||
"../../secret", # Deep traversal
|
||||
"", # Empty
|
||||
None # None
|
||||
]
|
||||
for path in invalid_paths:
|
||||
with self.subTest(path=path):
|
||||
self.assertFalse(validate_file_path(path), f"Should have failed for {path}")
|
||||
|
||||
def test_update_github_cdn_urls_rejects_invalid_repo(self):
|
||||
"""Test that update_github_cdn_urls rejects invalid repo before making requests."""
|
||||
repo = "user/repo/../malicious"
|
||||
success, message = update_github_cdn_urls(repo, "token", "file.md", [("f", "u")])
|
||||
|
||||
self.assertFalse(success)
|
||||
self.assertIn("Invalid GitHub repository format", message)
|
||||
|
||||
def test_update_github_cdn_urls_rejects_invalid_path(self):
|
||||
"""Test that update_github_cdn_urls rejects invalid path before making requests."""
|
||||
path = "../../../secret.txt"
|
||||
success, message = update_github_cdn_urls("user/repo", "token", path, [("f", "u")])
|
||||
|
||||
self.assertFalse(success)
|
||||
self.assertIn("Invalid file path", message)
|
||||
self.assertIn("Path traversal", message)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,173 @@
|
||||
import unittest
|
||||
import json
|
||||
import sys
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
import importlib
|
||||
|
||||
# Add project root to sys.path
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
|
||||
|
||||
# Clean up any potential pollution from other tests before we start
|
||||
if 'PIL' in sys.modules:
|
||||
# Check if it's a mock
|
||||
if isinstance(sys.modules['PIL'], MagicMock):
|
||||
del sys.modules['PIL']
|
||||
if 'PIL.PngImagePlugin' in sys.modules:
|
||||
del sys.modules['PIL.PngImagePlugin']
|
||||
if 'PIL.Image' in sys.modules:
|
||||
del sys.modules['PIL.Image']
|
||||
|
||||
# Now we can import real modules or mock them as we see fit LOCALLY
|
||||
# But wait, discord_image_node imports them at module level.
|
||||
# So we need to ensure environment is set up before importing it.
|
||||
|
||||
# Mock comfy modules
|
||||
sys.modules['comfy'] = MagicMock()
|
||||
sys.modules['comfy.cli_args'] = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args.disable_metadata = False
|
||||
sys.modules['comfy.utils'] = MagicMock()
|
||||
sys.modules['folder_paths'] = MagicMock()
|
||||
sys.modules['folder_paths'].get_output_directory = MagicMock(return_value="/tmp")
|
||||
sys.modules['folder_paths'].get_temp_directory = MagicMock(return_value="/tmp")
|
||||
sys.modules['folder_paths'].get_save_image_path = MagicMock(return_value=("/tmp", "test", 0, "", "test"))
|
||||
sys.modules['server'] = MagicMock()
|
||||
|
||||
# We need real PIL for this test to verify PngInfo
|
||||
try:
|
||||
import PIL.PngImagePlugin
|
||||
except ImportError:
|
||||
# If it failed because it was mocked out and we deleted it, reload
|
||||
pass
|
||||
|
||||
from nodes.image_node import DiscordSendSaveImage
|
||||
|
||||
# Check if torch is real or mocked
|
||||
try:
|
||||
import torch
|
||||
_torch_available = hasattr(torch, 'zeros') and callable(torch.zeros) and not isinstance(torch.zeros, MagicMock)
|
||||
except ImportError:
|
||||
_torch_available = False
|
||||
|
||||
class TestDiscordImageNodeOptimization(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.node = DiscordSendSaveImage()
|
||||
self.webhook_url = "https://discord.com/api/webhooks/12345/abcdef"
|
||||
self.github_token = "ghp_sensitive12345"
|
||||
|
||||
@unittest.skipUnless(_torch_available, "Test requires real torch for tensor iteration")
|
||||
def test_save_images_sanitization(self):
|
||||
# Create a mock image tensor using numpy (torch not available in CI)
|
||||
# The image_node iterates over images and accesses shape, so we need
|
||||
# an object that supports iteration and has proper shape
|
||||
import numpy as np
|
||||
|
||||
# Create a simple class that mimics torch.Tensor behavior for the node
|
||||
class MockTensor:
|
||||
def __init__(self, data):
|
||||
self._data = data
|
||||
self.shape = data.shape
|
||||
|
||||
def __len__(self):
|
||||
return len(self._data)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return self._data[idx]
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._data)
|
||||
|
||||
# Create a 1x64x64x3 "image batch" using numpy
|
||||
image_data = np.zeros((1, 64, 64, 3), dtype=np.float32)
|
||||
mock_image = MockTensor(image_data)
|
||||
|
||||
# Create prompt and extra_pnginfo with sensitive data
|
||||
prompt = {
|
||||
"3": {
|
||||
"inputs": {
|
||||
"webhook_url": self.webhook_url,
|
||||
"github_token": self.github_token,
|
||||
"seed": 123
|
||||
},
|
||||
"class_type": "DiscordSendSaveImage"
|
||||
}
|
||||
}
|
||||
|
||||
extra_pnginfo = {
|
||||
"workflow": {
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "DiscordSendSaveImage",
|
||||
"widgets_values": [self.webhook_url, "message", self.github_token]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
# Mock Image.save to check metadata
|
||||
# Also mock tensor_to_numpy_uint8 to bypass torch tensor conversion (torch is mocked)
|
||||
# and Image.fromarray to return a mock PIL Image with proper size attribute
|
||||
mock_pil_image = MagicMock()
|
||||
mock_pil_image.size = (64, 64)
|
||||
mock_pil_image.mode = 'RGB'
|
||||
|
||||
def mock_tensor_to_numpy(tensor):
|
||||
# Return a simple numpy-like array (64x64x3 zeros as uint8)
|
||||
import numpy as np
|
||||
return np.zeros((64, 64, 3), dtype=np.uint8)
|
||||
|
||||
with patch('PIL.Image.Image.save') as mock_save, \
|
||||
patch('nodes.image_node.tensor_to_numpy_uint8', side_effect=mock_tensor_to_numpy), \
|
||||
patch('PIL.Image.fromarray', return_value=mock_pil_image):
|
||||
self.node.save_images(
|
||||
images=mock_image,
|
||||
prompt=prompt,
|
||||
extra_pnginfo=extra_pnginfo,
|
||||
save_output=True,
|
||||
send_to_discord=False # Disable discord sending to focus on save/metadata
|
||||
)
|
||||
|
||||
# Check if save was called
|
||||
self.assertTrue(mock_save.called)
|
||||
|
||||
# Get the pnginfo passed to save
|
||||
args, kwargs = mock_save.call_args
|
||||
pnginfo = kwargs.get('pnginfo')
|
||||
self.assertIsNotNone(pnginfo)
|
||||
|
||||
found_prompt = False
|
||||
found_workflow = False
|
||||
|
||||
# Check chunks - PIL PngInfo internal structure
|
||||
for tag_type, data, after_idat in pnginfo.chunks:
|
||||
try:
|
||||
# decode data
|
||||
decoded = data.decode('latin-1')
|
||||
except Exception:
|
||||
# Skip chunks that can't be decoded
|
||||
continue
|
||||
|
||||
if '\0' in decoded:
|
||||
try:
|
||||
k, v = decoded.split('\0', 1)
|
||||
except ValueError:
|
||||
continue
|
||||
|
||||
if k == "prompt":
|
||||
found_prompt = True
|
||||
# Verify sensitive data is gone
|
||||
self.assertNotIn("discord.com/api/webhooks", v)
|
||||
self.assertNotIn("ghp_", v)
|
||||
|
||||
if k == "workflow":
|
||||
found_workflow = True
|
||||
self.assertNotIn("discord.com/api/webhooks", v)
|
||||
self.assertNotIn("ghp_", v)
|
||||
|
||||
self.assertTrue(found_prompt, "Prompt metadata not found")
|
||||
self.assertTrue(found_workflow, "Workflow metadata not found")
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+10
-4
@@ -3,7 +3,11 @@ import sys
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# Create a dummy torch module
|
||||
# IMPORTANT: Import real numpy FIRST before any mocking
|
||||
# This ensures test_power_of_two_math uses real numpy
|
||||
import numpy as np
|
||||
|
||||
# Create a dummy torch module (torch is not installed in CI)
|
||||
mock_torch = MagicMock()
|
||||
sys.modules["torch"] = mock_torch
|
||||
sys.modules["folder_paths"] = MagicMock()
|
||||
@@ -20,7 +24,7 @@ sys.modules["server"] = MagicMock()
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from discord_video_node import validate_video_for_discord
|
||||
from nodes.video_node import validate_video_for_discord
|
||||
|
||||
class TestPathLogic(unittest.TestCase):
|
||||
"""Tests for path and file validation logic."""
|
||||
@@ -65,10 +69,12 @@ class TestImageResizing(unittest.TestCase):
|
||||
|
||||
def test_power_of_two_math(self):
|
||||
"""Verify the power-of-two calculation logic used in the node."""
|
||||
import numpy as np
|
||||
# Use Python's built-in math module instead of numpy
|
||||
# to avoid test collection order issues with mocked modules
|
||||
import math
|
||||
|
||||
def calculate_nearest_pow2(dim):
|
||||
return 2 ** int(np.log2(dim) + 0.5)
|
||||
return 2 ** int(math.log2(dim) + 0.5)
|
||||
|
||||
self.assertEqual(calculate_nearest_pow2(500), 512)
|
||||
self.assertEqual(calculate_nearest_pow2(700), 512) # log2(700) = 9.45, +0.5 = 9.95, int=9, 2^9=512
|
||||
|
||||
@@ -0,0 +1,261 @@
|
||||
"""Tests for shared/discord/message_builder.py"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from shared.discord.message_builder import (
|
||||
build_metadata_section,
|
||||
build_prompt_section,
|
||||
build_discord_message,
|
||||
validate_message_content,
|
||||
format_file_info,
|
||||
format_file_size,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildMetadataSection(unittest.TestCase):
|
||||
"""Test build_metadata_section function."""
|
||||
|
||||
def test_empty_dict(self):
|
||||
"""Test with empty info dict returns empty string."""
|
||||
result = build_metadata_section({})
|
||||
self.assertEqual(result, "")
|
||||
|
||||
def test_with_date(self):
|
||||
"""Test metadata with date."""
|
||||
result = build_metadata_section({"date": "2026-01-20"})
|
||||
self.assertIn("**Date:** 2026-01-20", result)
|
||||
self.assertIn("**Information:**", result)
|
||||
|
||||
def test_with_time(self):
|
||||
"""Test metadata with time."""
|
||||
result = build_metadata_section({"time": "14-30-00"})
|
||||
self.assertIn("**Time:** 14-30-00", result)
|
||||
|
||||
def test_with_dimensions(self):
|
||||
"""Test metadata with dimensions."""
|
||||
result = build_metadata_section({"dimensions": "1920x1080"})
|
||||
self.assertIn("**Dimensions:** 1920x1080", result)
|
||||
|
||||
def test_with_format(self):
|
||||
"""Test metadata with file format."""
|
||||
result = build_metadata_section({}, file_format="png")
|
||||
self.assertIn("**Format:** PNG", result)
|
||||
|
||||
def test_with_frame_rate(self):
|
||||
"""Test metadata with frame rate."""
|
||||
result = build_metadata_section({}, frame_rate=30.0)
|
||||
self.assertIn("**Frame Rate:** 30.0 fps", result)
|
||||
|
||||
def test_custom_section_title(self):
|
||||
"""Test custom section title."""
|
||||
result = build_metadata_section({"date": "2026-01-20"}, section_title="Video Info")
|
||||
self.assertIn("**Video Info:**", result)
|
||||
|
||||
def test_exclude_options(self):
|
||||
"""Test excluding certain metadata."""
|
||||
info = {"date": "2026-01-20", "time": "14-30-00", "dimensions": "512x512"}
|
||||
result = build_metadata_section(info, include_date=False, include_time=False)
|
||||
self.assertNotIn("Date", result)
|
||||
self.assertNotIn("Time", result)
|
||||
self.assertIn("Dimensions", result)
|
||||
|
||||
def test_trailing_newline(self):
|
||||
"""Test that section ends with newline."""
|
||||
result = build_metadata_section({"date": "2026-01-20"})
|
||||
self.assertTrue(result.endswith("\n"))
|
||||
|
||||
|
||||
class TestBuildPromptSection(unittest.TestCase):
|
||||
"""Test build_prompt_section function."""
|
||||
|
||||
def test_no_prompts(self):
|
||||
"""Test with no prompts returns empty string."""
|
||||
result = build_prompt_section(None, None)
|
||||
self.assertEqual(result, "")
|
||||
|
||||
def test_empty_prompts(self):
|
||||
"""Test with empty prompts returns empty string."""
|
||||
result = build_prompt_section("", "")
|
||||
self.assertEqual(result, "")
|
||||
|
||||
def test_whitespace_prompts(self):
|
||||
"""Test with whitespace-only prompts returns empty string."""
|
||||
result = build_prompt_section(" ", " \n ")
|
||||
self.assertEqual(result, "")
|
||||
|
||||
def test_positive_only(self):
|
||||
"""Test with only positive prompt."""
|
||||
result = build_prompt_section("a beautiful sunset", None)
|
||||
self.assertIn("**Positive:**", result)
|
||||
self.assertIn("a beautiful sunset", result)
|
||||
self.assertNotIn("**Negative:**", result)
|
||||
|
||||
def test_negative_only(self):
|
||||
"""Test with only negative prompt."""
|
||||
result = build_prompt_section(None, "blurry, low quality")
|
||||
self.assertIn("**Negative:**", result)
|
||||
self.assertIn("blurry, low quality", result)
|
||||
self.assertNotIn("**Positive:**", result)
|
||||
|
||||
def test_both_prompts(self):
|
||||
"""Test with both prompts."""
|
||||
result = build_prompt_section("a cat", "dog")
|
||||
self.assertIn("**Positive:**", result)
|
||||
self.assertIn("a cat", result)
|
||||
self.assertIn("**Negative:**", result)
|
||||
self.assertIn("dog", result)
|
||||
|
||||
def test_custom_section_title(self):
|
||||
"""Test custom section title."""
|
||||
result = build_prompt_section("test", None, section_title="Custom Prompts")
|
||||
self.assertIn("**Custom Prompts:**", result)
|
||||
|
||||
def test_code_block_formatting(self):
|
||||
"""Test prompts are wrapped in code blocks."""
|
||||
result = build_prompt_section("test prompt", None)
|
||||
self.assertIn("```\ntest prompt\n```", result)
|
||||
|
||||
def test_non_string_conversion(self):
|
||||
"""Test that non-string prompts are converted."""
|
||||
result = build_prompt_section(12345, None)
|
||||
self.assertIn("12345", result)
|
||||
|
||||
|
||||
class TestBuildDiscordMessage(unittest.TestCase):
|
||||
"""Test build_discord_message function."""
|
||||
|
||||
def test_empty_message(self):
|
||||
"""Test building empty message."""
|
||||
result = build_discord_message()
|
||||
self.assertEqual(result, "")
|
||||
|
||||
def test_base_message_only(self):
|
||||
"""Test with just base message."""
|
||||
result = build_discord_message(base_message="Hello!")
|
||||
self.assertEqual(result, "Hello!")
|
||||
|
||||
def test_with_metadata(self):
|
||||
"""Test with metadata section."""
|
||||
result = build_discord_message(
|
||||
base_message="Image generated",
|
||||
metadata_section="\n**Info:** test"
|
||||
)
|
||||
self.assertIn("Image generated", result)
|
||||
self.assertIn("**Info:** test", result)
|
||||
|
||||
def test_with_all_sections(self):
|
||||
"""Test with all sections."""
|
||||
result = build_discord_message(
|
||||
base_message="Base",
|
||||
metadata_section="\nMeta",
|
||||
prompt_section="\nPrompt",
|
||||
additional_sections=["\nExtra1", "\nExtra2"]
|
||||
)
|
||||
self.assertIn("Base", result)
|
||||
self.assertIn("Meta", result)
|
||||
self.assertIn("Prompt", result)
|
||||
self.assertIn("Extra1", result)
|
||||
self.assertIn("Extra2", result)
|
||||
|
||||
def test_truncation(self):
|
||||
"""Test message truncation at max length."""
|
||||
long_message = "x" * 2500
|
||||
result = build_discord_message(base_message=long_message, max_length=2000)
|
||||
self.assertLessEqual(len(result), 2000)
|
||||
self.assertIn("[Message truncated]", result)
|
||||
|
||||
def test_no_truncation_under_limit(self):
|
||||
"""Test message not truncated when under limit."""
|
||||
message = "x" * 100
|
||||
result = build_discord_message(base_message=message)
|
||||
self.assertNotIn("truncated", result)
|
||||
|
||||
|
||||
class TestValidateMessageContent(unittest.TestCase):
|
||||
"""Test validate_message_content function."""
|
||||
|
||||
def test_empty_message(self):
|
||||
"""Test empty message is valid."""
|
||||
is_valid, msg = validate_message_content("")
|
||||
self.assertTrue(is_valid)
|
||||
self.assertIn("Empty message", msg)
|
||||
|
||||
def test_normal_message(self):
|
||||
"""Test normal message is valid."""
|
||||
is_valid, msg = validate_message_content("Hello world")
|
||||
self.assertTrue(is_valid)
|
||||
|
||||
def test_too_long_message(self):
|
||||
"""Test message over 2000 chars is invalid."""
|
||||
is_valid, msg = validate_message_content("x" * 2001)
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("2000 character limit", msg)
|
||||
|
||||
def test_message_with_prompts_section(self):
|
||||
"""Test message with Generation Prompts section."""
|
||||
message = "Test\n**Generation Prompts:**\nContent"
|
||||
is_valid, msg = validate_message_content(message)
|
||||
self.assertTrue(is_valid)
|
||||
self.assertNotIn("WARNING", msg)
|
||||
|
||||
def test_message_without_prompts_section(self):
|
||||
"""Test message without Generation Prompts section shows warning."""
|
||||
is_valid, msg = validate_message_content("Test message")
|
||||
self.assertTrue(is_valid)
|
||||
self.assertIn("WARNING", msg)
|
||||
|
||||
|
||||
class TestFormatFileSize(unittest.TestCase):
|
||||
"""Test format_file_size function."""
|
||||
|
||||
def test_bytes(self):
|
||||
"""Test formatting bytes."""
|
||||
self.assertEqual(format_file_size(500), "500 bytes")
|
||||
|
||||
def test_kilobytes(self):
|
||||
"""Test formatting kilobytes."""
|
||||
self.assertEqual(format_file_size(2048), "2.0 KB")
|
||||
|
||||
def test_megabytes(self):
|
||||
"""Test formatting megabytes."""
|
||||
self.assertEqual(format_file_size(5 * 1024 * 1024), "5.0 MB")
|
||||
|
||||
def test_gigabytes(self):
|
||||
"""Test formatting gigabytes."""
|
||||
self.assertEqual(format_file_size(2 * 1024 * 1024 * 1024), "2.00 GB")
|
||||
|
||||
def test_zero(self):
|
||||
"""Test formatting zero bytes."""
|
||||
self.assertEqual(format_file_size(0), "0 bytes")
|
||||
|
||||
|
||||
class TestFormatFileInfo(unittest.TestCase):
|
||||
"""Test format_file_info function."""
|
||||
|
||||
def test_basic_info(self):
|
||||
"""Test basic file info formatting."""
|
||||
result = format_file_info("image.png", 1024)
|
||||
self.assertIn("image.png", result)
|
||||
self.assertIn("1.0 KB", result)
|
||||
|
||||
def test_with_mime_type(self):
|
||||
"""Test file info with MIME type."""
|
||||
result = format_file_info("video.mp4", 1024 * 1024, "video/mp4")
|
||||
self.assertIn("video.mp4", result)
|
||||
self.assertIn("1.0 MB", result)
|
||||
self.assertIn("[video/mp4]", result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,110 @@
|
||||
import unittest
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Ensure we get real numpy, not a mock from other test files
|
||||
# Remove any mocked numpy before importing
|
||||
if 'numpy' in sys.modules and hasattr(sys.modules['numpy'], '_mock_name'):
|
||||
del sys.modules['numpy']
|
||||
|
||||
import numpy as np
|
||||
|
||||
# Verify numpy is real
|
||||
assert hasattr(np, 'arange'), "numpy.arange not found - numpy may be mocked"
|
||||
|
||||
class TestNumpyToSubprocess(unittest.TestCase):
|
||||
"""
|
||||
Verify that subprocess.Popen.stdin.write accepts numpy arrays directly.
|
||||
For subprocess.run(input=...), we need to be careful with numpy arrays due to ambiguity check in subprocess module.
|
||||
"""
|
||||
|
||||
def test_popen_stdin_write_numpy(self):
|
||||
"""Test writing numpy array to Popen.stdin"""
|
||||
# Create a small numpy array
|
||||
data = np.arange(256, dtype=np.uint8)
|
||||
|
||||
# Use python to echo input to output (cross-platform)
|
||||
cmd = [sys.executable, '-c', 'import sys; sys.stdout.buffer.write(sys.stdin.buffer.read())']
|
||||
p = subprocess.Popen(cmd, stdin=subprocess.PIPE, stdout=subprocess.PIPE)
|
||||
|
||||
# Write numpy array directly
|
||||
p.stdin.write(data)
|
||||
out, _ = p.communicate()
|
||||
|
||||
# Verify output matches input data bytes
|
||||
self.assertEqual(out, data.tobytes())
|
||||
self.assertEqual(len(out), 256)
|
||||
|
||||
def test_run_input_memoryview(self):
|
||||
"""
|
||||
Test passing numpy array as memoryview to subprocess.run input.
|
||||
"""
|
||||
data = np.arange(256, dtype=np.uint8)
|
||||
|
||||
# Use python to echo input to output (cross-platform)
|
||||
cmd = [sys.executable, '-c', 'import sys; sys.stdout.buffer.write(sys.stdin.buffer.read())']
|
||||
# memoryview works and avoids copy
|
||||
res = subprocess.run(cmd, input=memoryview(data), capture_output=True)
|
||||
|
||||
self.assertEqual(res.stdout, data.tobytes())
|
||||
self.assertEqual(len(res.stdout), 256)
|
||||
|
||||
def test_run_input_fixed_non_contiguous(self):
|
||||
"""
|
||||
Test that using ascontiguousarray makes the non-contiguous array accepted by subprocess.run
|
||||
"""
|
||||
# Create a 2D array and transpose it to make it non-contiguous
|
||||
data = np.zeros((10, 10), dtype=np.uint8)
|
||||
# Fill with some data
|
||||
for i in range(10):
|
||||
for j in range(10):
|
||||
data[i, j] = i + j
|
||||
|
||||
# Transpose creates a non-contiguous view
|
||||
transposed_data = data.T
|
||||
self.assertFalse(transposed_data.flags['C_CONTIGUOUS'])
|
||||
|
||||
# Fix it using ascontiguousarray
|
||||
contiguous_data = np.ascontiguousarray(transposed_data)
|
||||
self.assertTrue(contiguous_data.flags['C_CONTIGUOUS'])
|
||||
|
||||
# Now pass to subprocess
|
||||
mv = memoryview(contiguous_data)
|
||||
cmd = [sys.executable, '-c', 'import sys; sys.stdout.buffer.write(sys.stdin.buffer.read())']
|
||||
res = subprocess.run(cmd, input=mv, capture_output=True)
|
||||
self.assertEqual(res.stdout, contiguous_data.tobytes())
|
||||
|
||||
def test_popen_stdin_write_fixed_non_contiguous(self):
|
||||
"""
|
||||
Test writing fixed (made contiguous) numpy array to Popen.stdin.
|
||||
"""
|
||||
# Create a 2D array and transpose it to make it non-contiguous
|
||||
data = np.zeros((10, 10), dtype=np.uint8)
|
||||
# Fill with some data
|
||||
for i in range(10):
|
||||
for j in range(10):
|
||||
data[i, j] = i + j
|
||||
|
||||
transposed_data = data.T
|
||||
self.assertFalse(transposed_data.flags['C_CONTIGUOUS'])
|
||||
|
||||
# Fix it
|
||||
contiguous_data = np.ascontiguousarray(transposed_data)
|
||||
self.assertTrue(contiguous_data.flags['C_CONTIGUOUS'])
|
||||
|
||||
# Use python to echo input to output (cross-platform)
|
||||
cmd = [sys.executable, '-c', 'import sys; sys.stdout.buffer.write(sys.stdin.buffer.read())']
|
||||
p = subprocess.Popen(cmd, stdin=subprocess.PIPE, stdout=subprocess.PIPE)
|
||||
|
||||
try:
|
||||
# Should succeed now
|
||||
p.stdin.write(contiguous_data)
|
||||
out, _ = p.communicate()
|
||||
self.assertEqual(out, contiguous_data.tobytes())
|
||||
|
||||
except Exception as e:
|
||||
self.fail(f"Caught unexpected exception: {e}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,75 @@
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# Mock dependencies to allow testing in isolation
|
||||
sys.modules["folder_paths"] = MagicMock()
|
||||
sys.modules["comfy"] = MagicMock()
|
||||
sys.modules["server"] = MagicMock()
|
||||
sys.modules["requests"] = MagicMock()
|
||||
sys.modules["PIL"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["torch.nn"] = MagicMock()
|
||||
sys.modules["torch.nn.functional"] = MagicMock()
|
||||
|
||||
# Add project root to sys.path
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
from shared.path_utils import validate_path_is_safe
|
||||
|
||||
class TestPathSecurity(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.test_dir = os.path.join(os.getcwd(), "test_safe_env")
|
||||
os.makedirs(self.test_dir, exist_ok=True)
|
||||
|
||||
def tearDown(self):
|
||||
# Clean up would go here, but since we use temp dirs or mocks, it's fine.
|
||||
# Ideally use tempfile.TemporaryDirectory but this is simple.
|
||||
import shutil
|
||||
if os.path.exists(self.test_dir):
|
||||
shutil.rmtree(self.test_dir)
|
||||
|
||||
def test_absolute_path_blocked_with_base_dir(self):
|
||||
# This test checks if validate_path_is_safe blocks writing to /tmp
|
||||
# when base_dir is provided.
|
||||
|
||||
base_dir = self.test_dir
|
||||
test_path = "/tmp/sentinel_test_file.txt"
|
||||
|
||||
try:
|
||||
validate_path_is_safe(test_path, base_dir=base_dir)
|
||||
self.fail("Should have raised ValueError")
|
||||
except ValueError as e:
|
||||
self.assertIn("outside the allowed directory", str(e))
|
||||
|
||||
def test_traversal_blocked_with_base_dir(self):
|
||||
# Create a directory structure
|
||||
base_dir = self.test_dir
|
||||
subdir = os.path.join(base_dir, "subdir")
|
||||
os.makedirs(subdir, exist_ok=True)
|
||||
|
||||
# ../ traversal
|
||||
# This path resolves to outside base_dir
|
||||
test_path = os.path.abspath(os.path.join(subdir, "../../test_safe_escape.txt"))
|
||||
|
||||
try:
|
||||
validate_path_is_safe(test_path, base_dir=base_dir)
|
||||
self.fail("Should have raised ValueError")
|
||||
except ValueError as e:
|
||||
self.assertIn("outside the allowed directory", str(e))
|
||||
|
||||
def test_valid_path_allowed_with_base_dir(self):
|
||||
base_dir = self.test_dir
|
||||
test_path = os.path.join(base_dir, "valid_file.txt")
|
||||
|
||||
# Should not raise
|
||||
validate_path_is_safe(test_path, base_dir=base_dir)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,141 @@
|
||||
"""Tests for shared/path_utils.py"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
import tempfile
|
||||
import shutil
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from shared.path_utils import (
|
||||
get_output_directory,
|
||||
ensure_directory_exists,
|
||||
get_unique_filepath,
|
||||
)
|
||||
|
||||
|
||||
class TestGetOutputDirectory(unittest.TestCase):
|
||||
"""Test get_output_directory function."""
|
||||
|
||||
def setUp(self):
|
||||
"""Create temporary directories for testing."""
|
||||
self.test_dir = tempfile.mkdtemp()
|
||||
self.output_dir = os.path.join(self.test_dir, "output")
|
||||
self.temp_dir = os.path.join(self.test_dir, "temp")
|
||||
os.makedirs(self.output_dir)
|
||||
os.makedirs(self.temp_dir)
|
||||
|
||||
def tearDown(self):
|
||||
"""Clean up temporary directories."""
|
||||
shutil.rmtree(self.test_dir)
|
||||
|
||||
def test_save_output_true(self):
|
||||
"""Test output directory when saving is enabled."""
|
||||
result = get_output_directory(
|
||||
save_output=True,
|
||||
comfy_output_dir=self.output_dir,
|
||||
temp_dir=self.temp_dir,
|
||||
)
|
||||
expected = os.path.join(self.output_dir, "discord_output")
|
||||
self.assertEqual(result, expected)
|
||||
self.assertTrue(os.path.exists(result))
|
||||
|
||||
def test_save_output_false(self):
|
||||
"""Test temp directory when saving is disabled."""
|
||||
result = get_output_directory(
|
||||
save_output=False,
|
||||
comfy_output_dir=self.output_dir,
|
||||
temp_dir=self.temp_dir,
|
||||
)
|
||||
self.assertEqual(result, self.temp_dir)
|
||||
|
||||
def test_custom_subfolder(self):
|
||||
"""Test with custom subfolder name."""
|
||||
result = get_output_directory(
|
||||
save_output=True,
|
||||
comfy_output_dir=self.output_dir,
|
||||
temp_dir=self.temp_dir,
|
||||
subfolder="custom_folder",
|
||||
)
|
||||
expected = os.path.join(self.output_dir, "custom_folder")
|
||||
self.assertEqual(result, expected)
|
||||
self.assertTrue(os.path.exists(result))
|
||||
|
||||
|
||||
class TestEnsureDirectoryExists(unittest.TestCase):
|
||||
"""Test ensure_directory_exists function."""
|
||||
|
||||
def setUp(self):
|
||||
"""Create temporary directory for testing."""
|
||||
self.test_dir = tempfile.mkdtemp()
|
||||
|
||||
def tearDown(self):
|
||||
"""Clean up temporary directories."""
|
||||
shutil.rmtree(self.test_dir)
|
||||
|
||||
def test_creates_directory(self):
|
||||
"""Test that directory is created if it doesn't exist."""
|
||||
new_dir = os.path.join(self.test_dir, "new_directory")
|
||||
self.assertFalse(os.path.exists(new_dir))
|
||||
result = ensure_directory_exists(new_dir)
|
||||
self.assertTrue(os.path.exists(new_dir))
|
||||
self.assertEqual(result, new_dir)
|
||||
|
||||
def test_existing_directory(self):
|
||||
"""Test that existing directory is not affected."""
|
||||
result = ensure_directory_exists(self.test_dir)
|
||||
self.assertTrue(os.path.exists(self.test_dir))
|
||||
self.assertEqual(result, self.test_dir)
|
||||
|
||||
def test_nested_directories(self):
|
||||
"""Test creating nested directories."""
|
||||
nested = os.path.join(self.test_dir, "a", "b", "c")
|
||||
result = ensure_directory_exists(nested)
|
||||
self.assertTrue(os.path.exists(nested))
|
||||
self.assertEqual(result, nested)
|
||||
|
||||
|
||||
class TestGetUniqueFilepath(unittest.TestCase):
|
||||
"""Test get_unique_filepath function."""
|
||||
|
||||
def test_basic_filepath(self):
|
||||
"""Test basic filepath generation."""
|
||||
result = get_unique_filepath("/output", "image", ".png")
|
||||
self.assertEqual(result, "/output/image.png")
|
||||
|
||||
def test_with_counter(self):
|
||||
"""Test filepath with counter."""
|
||||
result = get_unique_filepath("/output", "image", ".png", counter=5)
|
||||
self.assertEqual(result, "/output/image_00005.png")
|
||||
|
||||
def test_counter_formatting(self):
|
||||
"""Test counter is formatted with leading zeros."""
|
||||
result = get_unique_filepath("/output", "image", ".jpg", counter=123)
|
||||
self.assertEqual(result, "/output/image_00123.jpg")
|
||||
|
||||
def test_extension_without_dot(self):
|
||||
"""Test extension is normalized if dot is missing."""
|
||||
result = get_unique_filepath("/output", "video", "mp4")
|
||||
self.assertEqual(result, "/output/video.mp4")
|
||||
|
||||
def test_extension_with_dot(self):
|
||||
"""Test extension with dot works correctly."""
|
||||
result = get_unique_filepath("/output", "video", ".mp4")
|
||||
self.assertEqual(result, "/output/video.mp4")
|
||||
|
||||
def test_counter_zero(self):
|
||||
"""Test counter value of zero."""
|
||||
result = get_unique_filepath("/output", "frame", ".png", counter=0)
|
||||
self.assertEqual(result, "/output/frame_00000.png")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,151 @@
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
import sys
|
||||
import os
|
||||
import numpy as np
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# Mock dependencies before importing nodes.video_node
|
||||
mock_torch = MagicMock()
|
||||
sys.modules["torch"] = mock_torch
|
||||
sys.modules["folder_paths"] = MagicMock()
|
||||
sys.modules["comfy"] = MagicMock()
|
||||
sys.modules["comfy.cli_args"] = MagicMock()
|
||||
sys.modules["comfy.utils"] = MagicMock()
|
||||
sys.modules["server"] = MagicMock()
|
||||
|
||||
# Mock PIL
|
||||
mock_pil = MagicMock()
|
||||
sys.modules["PIL"] = mock_pil
|
||||
sys.modules["PIL.Image"] = mock_pil
|
||||
sys.modules["PIL.PngImagePlugin"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Mock shared modules' submodules to avoid dependency issues
|
||||
sys.modules["shared.workflow"] = MagicMock()
|
||||
sys.modules["shared.workflow.sanitizer"] = MagicMock()
|
||||
sys.modules["shared.workflow.prompt_extractor"] = MagicMock()
|
||||
sys.modules["shared.workflow.workflow_builder"] = MagicMock()
|
||||
|
||||
sys.modules["shared.discord"] = MagicMock()
|
||||
sys.modules["shared.discord.webhook_client"] = MagicMock()
|
||||
sys.modules["shared.discord.message_builder"] = MagicMock()
|
||||
sys.modules["shared.discord.cdn_extractor"] = MagicMock()
|
||||
|
||||
sys.modules["shared.github_integration"] = MagicMock()
|
||||
sys.modules["shared.logging_config"] = MagicMock()
|
||||
sys.modules["shared.filename_utils"] = MagicMock()
|
||||
sys.modules["shared.path_utils"] = MagicMock()
|
||||
|
||||
# Mock shared.media siblings
|
||||
sys.modules["shared.media.format_utils"] = MagicMock()
|
||||
sys.modules["shared.media.video_encoder"] = MagicMock()
|
||||
|
||||
# Note: We do NOT mock "shared", "shared.media", or "shared.media.image_processing"
|
||||
# because we want to load the real code for testing.
|
||||
|
||||
# Define the function logic we want to verify (simulating the generator consumer)
|
||||
def consume_chunks(chunks):
|
||||
pil_images = []
|
||||
for chunk in chunks:
|
||||
# Proposed logic for nodes/video_node.py
|
||||
if len(chunk.shape) == 4:
|
||||
# Batched chunk (B, H, W, C)
|
||||
for i in range(chunk.shape[0]):
|
||||
pil_images.append(f"image_from_batch_{i}")
|
||||
else:
|
||||
# Single frame chunk (H, W, C)
|
||||
pil_images.append("image_from_single")
|
||||
return pil_images
|
||||
|
||||
class TestPILBatchOptimization(unittest.TestCase):
|
||||
|
||||
def test_consumer_logic_mixed_chunks(self):
|
||||
"""Test that the consumer logic correctly handles mixed 4D and 3D chunks."""
|
||||
|
||||
# 1. 4D Chunk (Batch of 2)
|
||||
chunk_batch = np.zeros((2, 10, 10, 3), dtype=np.uint8)
|
||||
|
||||
# 2. 3D Chunk (Single frame) - simulating what happens if generator yields single frame
|
||||
chunk_single = np.zeros((10, 10, 3), dtype=np.uint8)
|
||||
|
||||
chunks = [chunk_batch, chunk_single]
|
||||
|
||||
# Run consumer logic
|
||||
images = consume_chunks(chunks)
|
||||
|
||||
# Verify results
|
||||
# Should have 2 from batch + 1 from single = 3 images
|
||||
self.assertEqual(len(images), 3)
|
||||
self.assertEqual(images[0], "image_from_batch_0")
|
||||
self.assertEqual(images[1], "image_from_batch_1")
|
||||
self.assertEqual(images[2], "image_from_single")
|
||||
|
||||
def test_process_batched_images_integration(self):
|
||||
"""
|
||||
Verify that we can import and run process_batched_images with mocks,
|
||||
and that it chunks correctly.
|
||||
"""
|
||||
# Import needs to happen after mocks are set up
|
||||
from shared.media.image_processing import process_batched_images
|
||||
|
||||
# Setup mock tensor
|
||||
# We need to make sure isinstance(t, torch.Tensor) works
|
||||
|
||||
tensor_len = 5
|
||||
batch_size = 2
|
||||
|
||||
# Mock slicing
|
||||
def getitem(self, idx):
|
||||
# idx is a slice object
|
||||
start = idx.start
|
||||
stop = idx.stop
|
||||
if stop > tensor_len:
|
||||
stop = tensor_len
|
||||
size = stop - start
|
||||
return f"slice_{size}"
|
||||
|
||||
# Create a class with __len__ and __getitem__ defined
|
||||
class MockTensor:
|
||||
def __len__(self):
|
||||
return tensor_len
|
||||
def __getitem__(self, idx):
|
||||
return getitem(self, idx)
|
||||
|
||||
mock_torch.Tensor = MockTensor
|
||||
|
||||
mock_tensor = mock_torch.Tensor()
|
||||
|
||||
# Mock tensor_to_numpy_uint8 to return numpy arrays of appropriate shape
|
||||
# It needs to return (Size, H, W, C)
|
||||
with patch('shared.media.image_processing.tensor_to_numpy_uint8') as mock_t2n:
|
||||
def side_effect(slice_obj):
|
||||
# parse size from string "slice_N"
|
||||
size = int(slice_obj.split('_')[1])
|
||||
return np.zeros((size, 10, 10, 3), dtype=np.uint8)
|
||||
|
||||
mock_t2n.side_effect = side_effect
|
||||
|
||||
# Run generator
|
||||
generator = process_batched_images(mock_tensor, batch_size=batch_size)
|
||||
chunks = list(generator)
|
||||
|
||||
# Expected:
|
||||
# 5 items, batch 2
|
||||
# 1. Size 2
|
||||
# 2. Size 2
|
||||
# 3. Size 1
|
||||
|
||||
self.assertEqual(len(chunks), 3)
|
||||
self.assertEqual(chunks[0].shape[0], 2)
|
||||
self.assertEqual(chunks[1].shape[0], 2)
|
||||
self.assertEqual(chunks[2].shape[0], 1)
|
||||
|
||||
# Verify they are all 4D arrays (B, H, W, C)
|
||||
for c in chunks:
|
||||
self.assertEqual(len(c.shape), 4)
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
@@ -0,0 +1,125 @@
|
||||
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
sys.modules["PIL"] = MagicMock()
|
||||
sys.modules["PIL.Image"] = MagicMock()
|
||||
sys.modules["PIL.PngImagePlugin"] = MagicMock()
|
||||
sys.modules["comfy"] = MagicMock()
|
||||
sys.modules["comfy.cli_args"] = MagicMock()
|
||||
sys.modules["comfy.utils"] = MagicMock()
|
||||
sys.modules["server"] = MagicMock()
|
||||
|
||||
# Mock folder_paths
|
||||
mock_folder_paths = MagicMock()
|
||||
sys.modules["folder_paths"] = mock_folder_paths
|
||||
|
||||
# Add parent directory to path
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# Import the node
|
||||
from nodes.video_node import DiscordSendSaveVideo
|
||||
|
||||
class TestTempFileLeak(unittest.TestCase):
|
||||
def setUp(self):
|
||||
# Create a real temporary directory for our test
|
||||
self.test_dir = tempfile.mkdtemp()
|
||||
self.output_dir = os.path.join(self.test_dir, "output")
|
||||
self.temp_dir = os.path.join(self.test_dir, "temp")
|
||||
os.makedirs(self.output_dir)
|
||||
os.makedirs(self.temp_dir)
|
||||
|
||||
# Configure folder_paths mock
|
||||
mock_folder_paths.get_output_directory.return_value = self.output_dir
|
||||
mock_folder_paths.get_temp_directory.return_value = self.temp_dir
|
||||
|
||||
# Mock get_save_image_path to return predictable paths
|
||||
# full_output_folder, filename, counter, subfolder, filename_prefix
|
||||
mock_folder_paths.get_save_image_path.return_value = (
|
||||
self.output_dir, "ComfyUI-Video", 1, "", "ComfyUI-Video"
|
||||
)
|
||||
|
||||
# Instantiate the node
|
||||
self.node = DiscordSendSaveVideo()
|
||||
|
||||
# Create a dummy image tensor mock
|
||||
self.dummy_image = MagicMock()
|
||||
self.dummy_image.shape = (512, 512, 3) # height, width, channels
|
||||
|
||||
# Mock tensor_to_numpy_uint8 in discordsend_utils
|
||||
self.patcher_numpy = patch("nodes.video_node.tensor_to_numpy_uint8")
|
||||
self.mock_numpy_conv = self.patcher_numpy.start()
|
||||
# Return a dummy numpy array
|
||||
import numpy as np
|
||||
self.mock_numpy_conv.return_value = np.zeros((512, 512, 3), dtype=np.uint8)
|
||||
|
||||
def tearDown(self):
|
||||
self.patcher_numpy.stop()
|
||||
shutil.rmtree(self.test_dir)
|
||||
|
||||
@patch("nodes.video_node.subprocess.Popen")
|
||||
@patch("nodes.video_node.subprocess.run")
|
||||
@patch("nodes.base_node.send_to_discord_with_retry")
|
||||
@patch("nodes.video_node.Image")
|
||||
@patch("nodes.video_node.os.path.getsize")
|
||||
@patch("nodes.video_node.validate_video_for_discord")
|
||||
def test_temp_file_leak(self, mock_validate, mock_getsize, mock_image, mock_send, mock_run, mock_popen):
|
||||
# Setup mocks
|
||||
mock_process = MagicMock()
|
||||
mock_process.returncode = 0
|
||||
mock_process.stdin = MagicMock()
|
||||
mock_popen.return_value = mock_process
|
||||
|
||||
mock_send.return_value.status_code = 200
|
||||
mock_send.return_value.json.return_value = {}
|
||||
|
||||
mock_getsize.return_value = 1024 * 1024 # 1MB
|
||||
mock_validate.return_value = (True, "Valid")
|
||||
|
||||
# Create a fake output file that subprocess would have created
|
||||
fake_output_path = os.path.join(self.output_dir, "ComfyUI-Video_00001.mp4")
|
||||
with open(fake_output_path, "wb") as f:
|
||||
f.write(b"fake video content")
|
||||
|
||||
# Mock subprocess.run to simulate creation of optimized file
|
||||
def side_effect_run(args, **kwargs):
|
||||
# The last argument is the output file path
|
||||
output_file = args[-1]
|
||||
if "discord_optimized_" in output_file:
|
||||
# Create the file
|
||||
with open(output_file, "wb") as f:
|
||||
f.write(b"optimized video content")
|
||||
return MagicMock(returncode=0)
|
||||
|
||||
mock_run.side_effect = side_effect_run
|
||||
|
||||
# Run the node
|
||||
self.node.save_video(
|
||||
images=[self.dummy_image],
|
||||
send_to_discord=True,
|
||||
webhook_url="https://discord.com/api/webhooks/123/abc",
|
||||
format="video/h264-mp4",
|
||||
save_output=True
|
||||
)
|
||||
|
||||
# Check if any file in temp_dir contains "discord_optimized_"
|
||||
temp_files = os.listdir(self.temp_dir)
|
||||
optimized_files = [f for f in temp_files if "discord_optimized_" in f]
|
||||
|
||||
print(f"Files remaining in temp dir: {temp_files}")
|
||||
|
||||
# This assertion is expected to FAIL if the leak exists (because we want 0 files)
|
||||
# Or pass if we assert > 0 to prove the leak.
|
||||
# After fix, we expect 0 files.
|
||||
self.assertEqual(len(optimized_files), 0, "Temporary optimized file should have been cleaned up")
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,246 @@
|
||||
import unittest
|
||||
import sys
|
||||
import os
|
||||
import tempfile
|
||||
import shutil
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# Add project root to sys.path
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
|
||||
|
||||
# Mock comfy modules
|
||||
if 'comfy' not in sys.modules:
|
||||
sys.modules['comfy'] = MagicMock()
|
||||
if 'comfy.cli_args' not in sys.modules:
|
||||
sys.modules['comfy.cli_args'] = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args.disable_metadata = False
|
||||
if 'comfy.utils' not in sys.modules:
|
||||
sys.modules['comfy.utils'] = MagicMock()
|
||||
if 'folder_paths' not in sys.modules:
|
||||
sys.modules['folder_paths'] = MagicMock()
|
||||
if 'server' not in sys.modules:
|
||||
sys.modules['server'] = MagicMock()
|
||||
|
||||
# Mock heavy/external dependencies
|
||||
sys.modules['torch'] = MagicMock()
|
||||
sys.modules['cv2'] = MagicMock()
|
||||
# We rely on real Pillow and numpy being installed and used
|
||||
|
||||
# Import the node
|
||||
try:
|
||||
from nodes.video_node import DiscordSendSaveVideo
|
||||
except ImportError:
|
||||
raise
|
||||
|
||||
class TestSymlinkAttack(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.test_dir = tempfile.mkdtemp()
|
||||
self.output_dir = os.path.join(self.test_dir, "output")
|
||||
self.temp_dir = os.path.join(self.test_dir, "temp")
|
||||
os.makedirs(self.output_dir)
|
||||
os.makedirs(self.temp_dir)
|
||||
|
||||
self.node = DiscordSendSaveVideo()
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.test_dir)
|
||||
|
||||
def test_overwrite_symlink_vulnerability(self):
|
||||
"""Test overwriting a direct symlink file."""
|
||||
# Create a target file
|
||||
target_file = os.path.join(self.test_dir, "target.txt")
|
||||
with open(target_file, "w") as f:
|
||||
f.write("Original Content")
|
||||
|
||||
# Create symlinks in output dir pointing to target file
|
||||
symlink_video = os.path.join(self.output_dir, "ComfyUI-Video_00001.mp4")
|
||||
os.symlink(target_file, symlink_video)
|
||||
|
||||
# Patch folder_paths on the node module
|
||||
with patch('nodes.video_node.folder_paths') as mock_folder_paths:
|
||||
mock_folder_paths.get_save_image_path.return_value = (
|
||||
self.output_dir,
|
||||
"ComfyUI-Video",
|
||||
1,
|
||||
"",
|
||||
"ComfyUI-Video"
|
||||
)
|
||||
|
||||
# Mock images input
|
||||
mock_images = [MagicMock()]
|
||||
mock_images[0].shape = (64, 64, 3)
|
||||
|
||||
with patch('nodes.video_node.tensor_to_numpy_uint8') as mock_t2n:
|
||||
import numpy as np
|
||||
mock_t2n.return_value = np.zeros((64, 64, 3), dtype=np.uint8)
|
||||
|
||||
with patch('nodes.video_node.ffmpeg_path', "ffmpeg"):
|
||||
with patch('subprocess.Popen') as mock_popen:
|
||||
mock_popen.return_value = MagicMock()
|
||||
|
||||
# Expect ValueError when trying to overwrite symlink
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self.node.save_video(
|
||||
images=mock_images,
|
||||
overwrite_last=True,
|
||||
format="video/h264-mp4",
|
||||
save_output=True,
|
||||
frame_rate=1.0
|
||||
)
|
||||
|
||||
self.assertIn("symlink", str(context.exception))
|
||||
print("SUCCESS: Direct symlink overwrite prevented.")
|
||||
|
||||
def test_parent_directory_symlink_vulnerability(self):
|
||||
"""Test writing to a path where a parent directory is a symlink."""
|
||||
# Create a real directory outside the intended output
|
||||
secret_dir = os.path.join(self.test_dir, "secret")
|
||||
os.makedirs(secret_dir)
|
||||
|
||||
# Create a symlink inside output_dir pointing to secret_dir
|
||||
# /test_dir/output/evil_link -> /test_dir/secret
|
||||
evil_link = os.path.join(self.output_dir, "evil_link")
|
||||
os.symlink(secret_dir, evil_link)
|
||||
|
||||
# We want to write to /test_dir/output/evil_link/file.mp4
|
||||
# which resolves to /test_dir/secret/file.mp4
|
||||
|
||||
# Patch folder_paths to return the evil_link as the directory
|
||||
with patch('nodes.video_node.folder_paths') as mock_folder_paths:
|
||||
mock_folder_paths.get_save_image_path.return_value = (
|
||||
evil_link,
|
||||
"ComfyUI-Video",
|
||||
1,
|
||||
"",
|
||||
"ComfyUI-Video"
|
||||
)
|
||||
|
||||
mock_images = [MagicMock()]
|
||||
mock_images[0].shape = (64, 64, 3)
|
||||
|
||||
with patch('nodes.video_node.tensor_to_numpy_uint8') as mock_t2n:
|
||||
import numpy as np
|
||||
mock_t2n.return_value = np.zeros((64, 64, 3), dtype=np.uint8)
|
||||
|
||||
with patch('nodes.video_node.ffmpeg_path', "ffmpeg"):
|
||||
with patch('subprocess.Popen') as mock_popen:
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self.node.save_video(
|
||||
images=mock_images,
|
||||
overwrite_last=True,
|
||||
format="video/h264-mp4",
|
||||
save_output=True,
|
||||
frame_rate=1.0
|
||||
)
|
||||
|
||||
# Check for either the parent dir message or the generic mismatch message
|
||||
error_msg = str(context.exception)
|
||||
self.assertTrue(
|
||||
"Writing through directory symlinks is not allowed" in error_msg or
|
||||
"Symlinks in output paths are not allowed" in error_msg,
|
||||
f"Unexpected error message: {error_msg}"
|
||||
)
|
||||
print("SUCCESS: Parent directory symlink prevented.")
|
||||
|
||||
def test_non_existent_directory_symlink_bypass(self):
|
||||
"""Test where intermediate directory doesn't exist but parent is symlink."""
|
||||
# /test_dir/secret
|
||||
secret_dir = os.path.join(self.test_dir, "secret")
|
||||
os.makedirs(secret_dir)
|
||||
|
||||
# /test_dir/output/link -> /test_dir/secret
|
||||
link_dir = os.path.join(self.output_dir, "link")
|
||||
os.symlink(secret_dir, link_dir)
|
||||
|
||||
# Target: /test_dir/output/link/subdir/file.mp4
|
||||
# 'subdir' does not exist yet.
|
||||
target_dir = os.path.join(link_dir, "subdir")
|
||||
# Do NOT create target_dir.
|
||||
|
||||
# Patch folder_paths to return the non-existent target_dir
|
||||
with patch('nodes.video_node.folder_paths') as mock_folder_paths:
|
||||
mock_folder_paths.get_save_image_path.return_value = (
|
||||
target_dir,
|
||||
"ComfyUI-Video",
|
||||
1,
|
||||
"",
|
||||
"ComfyUI-Video"
|
||||
)
|
||||
|
||||
# Use os.makedirs real implementation to create the directory if the node calls it
|
||||
# But here we assume the validation happens before directory creation or during path validation
|
||||
|
||||
mock_images = [MagicMock()]
|
||||
mock_images[0].shape = (64, 64, 3)
|
||||
|
||||
with patch('nodes.video_node.tensor_to_numpy_uint8') as mock_t2n:
|
||||
import numpy as np
|
||||
mock_t2n.return_value = np.zeros((64, 64, 3), dtype=np.uint8)
|
||||
|
||||
with patch('nodes.video_node.ffmpeg_path', "ffmpeg"):
|
||||
with patch('subprocess.Popen') as mock_popen:
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self.node.save_video(
|
||||
images=mock_images,
|
||||
overwrite_last=True,
|
||||
format="video/h264-mp4",
|
||||
save_output=True,
|
||||
frame_rate=1.0
|
||||
)
|
||||
|
||||
error_msg = str(context.exception)
|
||||
self.assertTrue(
|
||||
"Symlinks in output paths are not allowed" in error_msg or
|
||||
"Path component" in error_msg and "is a symlink" in error_msg,
|
||||
f"Unexpected error message: {error_msg}"
|
||||
)
|
||||
print("SUCCESS: Non-existent directory symlink bypass prevented.")
|
||||
|
||||
def test_vhs_format_bypass(self):
|
||||
"""Test that VHS format path recalculation is also validated."""
|
||||
# This test tries to exploit the path where 'is_vhs_format' is True
|
||||
# which changes the file extension and potentially bypasses early checks
|
||||
|
||||
target_file = os.path.join(self.test_dir, "target.mkv")
|
||||
with open(target_file, "w") as f:
|
||||
f.write("Original Content")
|
||||
|
||||
# Create symlink with different extension (mkv) that VHS might use
|
||||
symlink_video = os.path.join(self.output_dir, "ComfyUI-Video_00001.mkv")
|
||||
os.symlink(target_file, symlink_video)
|
||||
|
||||
with patch('nodes.video_node.folder_paths') as mock_folder_paths:
|
||||
mock_folder_paths.get_save_image_path.return_value = (
|
||||
self.output_dir,
|
||||
"ComfyUI-Video",
|
||||
1,
|
||||
"",
|
||||
"ComfyUI-Video"
|
||||
)
|
||||
|
||||
mock_images = [MagicMock()]
|
||||
mock_images[0].shape = (64, 64, 3)
|
||||
|
||||
with patch('nodes.video_node.tensor_to_numpy_uint8') as mock_t2n:
|
||||
import numpy as np
|
||||
mock_t2n.return_value = np.zeros((64, 64, 3), dtype=np.uint8)
|
||||
|
||||
with patch('nodes.video_node.ffmpeg_path', "ffmpeg"):
|
||||
# Mock has_vhs_formats to be True
|
||||
with patch('nodes.video_node.has_vhs_formats', True):
|
||||
with patch('subprocess.Popen') as mock_popen:
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self.node.save_video(
|
||||
images=mock_images,
|
||||
overwrite_last=True,
|
||||
format="video/mkv", # Custom format triggers VHS path
|
||||
save_output=True,
|
||||
frame_rate=1.0
|
||||
)
|
||||
|
||||
self.assertIn("symlink", str(context.exception))
|
||||
print("SUCCESS: VHS format path bypass prevented.")
|
||||
+94
-4
@@ -8,13 +8,21 @@ Or without pytest: python tests/test_utils.py
|
||||
import sys
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
import requests
|
||||
|
||||
# Mock dependencies before importing project modules
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
# sys.modules["PIL"] = MagicMock() # PIL might be installed, so maybe not mock it if not needed, but safer to mock if we don't rely on it for these tests
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from discordsend_utils.sanitizer import sanitize_json_for_export
|
||||
from discordsend_utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging, send_to_discord_with_retry
|
||||
from unittest.mock import patch, MagicMock
|
||||
from shared.workflow.sanitizer import sanitize_json_for_export
|
||||
from shared.discord.webhook_client import validate_webhook_url, sanitize_webhook_for_logging, send_to_discord_with_retry, DiscordWebhookClient
|
||||
from shared.github_integration import update_github_cdn_urls
|
||||
|
||||
|
||||
class TestSanitizer(unittest.TestCase):
|
||||
@@ -93,6 +101,31 @@ class TestSanitizer(unittest.TestCase):
|
||||
result = sanitize_json_for_export(test_data)
|
||||
self.assertEqual(result, test_data)
|
||||
|
||||
def test_sanitize_node_other_properties(self):
|
||||
"""Should sanitize other properties in nodes that are not inputs/widgets."""
|
||||
test_data = {
|
||||
"nodes": [
|
||||
{
|
||||
"id": 1,
|
||||
"type": "SomeNode",
|
||||
"inputs": {},
|
||||
"widgets_values": [],
|
||||
"extra": {
|
||||
"webhook_url": "https://discord.com/api/webhooks/123/abc"
|
||||
},
|
||||
"properties": {
|
||||
"nested": {
|
||||
"token": "ghp_secret"
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
result = sanitize_json_for_export(test_data)
|
||||
node = result["nodes"][0]
|
||||
self.assertEqual(node["extra"]["webhook_url"], "")
|
||||
self.assertEqual(node["properties"]["nested"]["token"], "")
|
||||
|
||||
|
||||
class TestWebhookValidation(unittest.TestCase):
|
||||
"""Tests for webhook URL validation."""
|
||||
@@ -131,6 +164,12 @@ class TestWebhookValidation(unittest.TestCase):
|
||||
is_valid, message = validate_webhook_url("http://localhost:8080/admin")
|
||||
self.assertFalse(is_valid)
|
||||
|
||||
def test_http_url_rejected(self):
|
||||
"""Should reject HTTP URLs (must be HTTPS)."""
|
||||
is_valid, message = validate_webhook_url("http://discord.com/api/webhooks/123/abc")
|
||||
self.assertFalse(is_valid)
|
||||
self.assertIn("must start with https://", message)
|
||||
|
||||
def test_ip_encoding_urls(self):
|
||||
"""Should reject alternate IP encodings."""
|
||||
self.assertFalse(validate_webhook_url("http://127.0.0.1")[0])
|
||||
@@ -162,7 +201,7 @@ class TestSSRFPrevention(unittest.TestCase):
|
||||
|
||||
self.assertIn("Invalid webhook URL", str(cm.exception))
|
||||
|
||||
@patch('discordsend_utils.discord_api.requests.post')
|
||||
@patch('shared.discord.webhook_client.requests.post')
|
||||
def test_send_to_discord_allows_valid_url(self, mock_post):
|
||||
"""Should allow valid Discord URLs."""
|
||||
valid_url = "https://discord.com/api/webhooks/123/abc"
|
||||
@@ -193,6 +232,57 @@ class TestWebhookSanitization(unittest.TestCase):
|
||||
self.assertEqual(result, "")
|
||||
|
||||
|
||||
class TestDiscordWebhookClient(unittest.TestCase):
|
||||
"""Tests for DiscordWebhookClient security features."""
|
||||
|
||||
@patch('shared.discord.webhook_client.requests.post')
|
||||
def test_exception_token_leakage(self, mock_post):
|
||||
"""Should redact tokens from exception messages in last_error."""
|
||||
token = "SUPER_SECRET_TOKEN"
|
||||
url = f"https://discord.com/api/webhooks/123456/{token}"
|
||||
client = DiscordWebhookClient(url)
|
||||
|
||||
# Configure mock to raise an exception containing the token
|
||||
error_message = f"Max retries exceeded with url: /api/webhooks/123456/{token}"
|
||||
mock_post.side_effect = requests.exceptions.ConnectionError(error_message)
|
||||
|
||||
success, result = client.send_message("Test message")
|
||||
|
||||
self.assertFalse(success)
|
||||
self.assertIn("error", result)
|
||||
self.assertNotIn(token, result["error"])
|
||||
self.assertIn("[REDACTED]", result["error"])
|
||||
|
||||
|
||||
class TestGitHubIntegration(unittest.TestCase):
|
||||
"""Tests for GitHub integration security features."""
|
||||
|
||||
@patch('shared.github_integration.requests.put')
|
||||
@patch('shared.github_integration.requests.get')
|
||||
def test_github_token_redaction_in_response(self, mock_get, mock_put):
|
||||
"""Should redact GitHub token from error messages including response text."""
|
||||
token = "ghp_SECRET_TOKEN"
|
||||
repo = "user/repo"
|
||||
file_path = "cdn_urls.md"
|
||||
|
||||
# Mock GET to return 404 (file doesn't exist)
|
||||
mock_get_response = MagicMock()
|
||||
mock_get_response.status_code = 404
|
||||
mock_get.return_value = mock_get_response
|
||||
|
||||
# Mock PUT to fail and return the token in the response text (simulating leak)
|
||||
mock_put_response = MagicMock()
|
||||
mock_put_response.status_code = 401
|
||||
mock_put_response.text = f"Bad credentials: {token} is invalid"
|
||||
mock_put.return_value = mock_put_response
|
||||
|
||||
success, message = update_github_cdn_urls(repo, token, file_path, [("test.png", "http://url")])
|
||||
|
||||
self.assertFalse(success)
|
||||
self.assertNotIn(token, message)
|
||||
self.assertIn("[REDACTED_TOKEN]", message)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests
|
||||
print("Running ComfyUI-DiscordSend utility tests...\n")
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
import unittest
|
||||
import sys
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# Add project root to sys.path
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
|
||||
|
||||
# Mock comfy modules needed for import
|
||||
sys.modules['comfy'] = MagicMock()
|
||||
sys.modules['comfy.cli_args'] = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args = MagicMock()
|
||||
sys.modules['comfy.cli_args'].args.disable_metadata = False
|
||||
sys.modules['comfy.utils'] = MagicMock()
|
||||
sys.modules['folder_paths'] = MagicMock()
|
||||
sys.modules['server'] = MagicMock()
|
||||
|
||||
# Mock torch if not available (video_node imports it)
|
||||
if 'torch' not in sys.modules:
|
||||
sys.modules['torch'] = MagicMock()
|
||||
|
||||
from nodes.image_node import DiscordSendSaveImage
|
||||
from nodes.video_node import DiscordSendSaveVideo
|
||||
|
||||
class TestUXTooltips(unittest.TestCase):
|
||||
def test_video_node_add_time_tooltip(self):
|
||||
"""Test that the add_time tooltip in video node contains the critical warning."""
|
||||
input_types = DiscordSendSaveVideo.INPUT_TYPES()
|
||||
add_time_config = input_types["optional"]["add_time"]
|
||||
tooltip = add_time_config[1]["tooltip"]
|
||||
|
||||
# Verify it HAS the warning
|
||||
self.assertIn("CRITICAL", tooltip)
|
||||
self.assertIn("single-frame playback", tooltip)
|
||||
self.assertIn("Add time", tooltip)
|
||||
|
||||
def test_image_node_add_time_tooltip(self):
|
||||
"""Test that the add_time tooltip in image node is standard."""
|
||||
input_types = DiscordSendSaveImage.INPUT_TYPES()
|
||||
add_time_config = input_types["optional"]["add_time"]
|
||||
tooltip = add_time_config[1]["tooltip"]
|
||||
|
||||
expected = "Add time (HH-MM-SS) to the filename."
|
||||
self.assertEqual(tooltip, expected)
|
||||
|
||||
def test_github_token_tooltip(self):
|
||||
"""Test that the github_token tooltip contains helpful instructions."""
|
||||
# Both nodes inherit from BaseDiscordNode, so check one
|
||||
input_types = DiscordSendSaveImage.INPUT_TYPES()
|
||||
token_config = input_types["optional"]["github_token"]
|
||||
tooltip = token_config[1]["tooltip"]
|
||||
|
||||
self.assertIn("Settings > Developer settings > Tokens", tooltip)
|
||||
self.assertIn("Requires 'repo' scope", tooltip)
|
||||
|
||||
def test_resize_method_clarity(self):
|
||||
"""Test that resize_method tooltip clarifies dependency on resize_to_power_of_2."""
|
||||
input_types = DiscordSendSaveImage.INPUT_TYPES()
|
||||
resize_config = input_types["optional"]["resize_method"]
|
||||
tooltip = resize_config[1]["tooltip"]
|
||||
|
||||
self.assertIn("ONLY when 'resize_to_power_of_2' is enabled", tooltip)
|
||||
self.assertIn("Ignored otherwise", tooltip)
|
||||
self.assertIn("lanczos: Best for photos", tooltip)
|
||||
|
||||
def test_overwrite_safety_warning(self):
|
||||
"""Test that overwrite_last tooltip contains safety warning in both nodes."""
|
||||
# Test Image Node
|
||||
input_types_img = DiscordSendSaveImage.INPUT_TYPES()
|
||||
tooltip_img = input_types_img["required"]["overwrite_last"][1]["tooltip"]
|
||||
|
||||
self.assertIn("CAUTION", tooltip_img)
|
||||
self.assertIn("REPLACE the previous file", tooltip_img)
|
||||
self.assertIn("dangerous for batch production", tooltip_img)
|
||||
self.assertIn("disable 'add_time' and 'add_date'", tooltip_img)
|
||||
|
||||
# Test Video Node
|
||||
input_types_vid = DiscordSendSaveVideo.INPUT_TYPES()
|
||||
tooltip_vid = input_types_vid["required"]["overwrite_last"][1]["tooltip"]
|
||||
|
||||
self.assertIn("CAUTION", tooltip_vid)
|
||||
self.assertIn("REPLACE the previous file", tooltip_vid)
|
||||
self.assertIn("dangerous for batch production", tooltip_vid)
|
||||
self.assertIn("Disabling 'add_time' to overwrite files will cause single-frame playback issues", tooltip_vid)
|
||||
|
||||
def test_include_video_info_tooltip(self):
|
||||
"""Test that include_video_info tooltip provides the helpful tip about add_time."""
|
||||
input_types = DiscordSendSaveVideo.INPUT_TYPES()
|
||||
info_config = input_types["optional"]["include_video_info"]
|
||||
tooltip = info_config[1]["tooltip"]
|
||||
|
||||
self.assertIn("TIP", tooltip)
|
||||
self.assertIn("Disable this instead of 'add_time'", tooltip)
|
||||
self.assertIn("avoid the Discord single-frame bug", tooltip)
|
||||
|
||||
def test_resize_to_power_of_2_tooltip(self):
|
||||
"""Test that resize_to_power_of_2 tooltip warns about aspect ratio distortion."""
|
||||
input_types = DiscordSendSaveImage.INPUT_TYPES()
|
||||
resize_config = input_types["optional"]["resize_to_power_of_2"]
|
||||
tooltip = resize_config[1]["tooltip"]
|
||||
|
||||
self.assertIn("May distort aspect ratio", tooltip)
|
||||
self.assertIn("Uses the algorithm selected in 'resize_method'", tooltip)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,70 @@
|
||||
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Mock modules that might be missing in the environment or causing issues
|
||||
sys.modules["folder_paths"] = MagicMock()
|
||||
sys.modules["comfy"] = MagicMock()
|
||||
sys.modules["server"] = MagicMock()
|
||||
sys.modules["torch"] = MagicMock()
|
||||
sys.modules["numpy"] = MagicMock()
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
|
||||
# Import the module under test
|
||||
sys.path.append(os.getcwd())
|
||||
|
||||
from shared.discord.webhook_client import DiscordWebhookClient, send_to_discord_with_retry, sanitize_token_from_text
|
||||
|
||||
class TestWebhookSecurity(unittest.TestCase):
|
||||
def test_token_leak_in_client_error(self):
|
||||
"""
|
||||
Test that webhook tokens are NOT leaked in client error details.
|
||||
"""
|
||||
webhook_url = "https://discord.com/api/webhooks/123456789/SuperSecretToken123"
|
||||
client = DiscordWebhookClient(webhook_url)
|
||||
|
||||
# Mock response to simulate a 400 error that echoes the URL
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 400
|
||||
# Simulate an API that echoes the request URL in the error body
|
||||
mock_response.text = f"Error processing request to {webhook_url}: Invalid payload"
|
||||
mock_response.content = mock_response.text.encode('utf-8')
|
||||
|
||||
with patch('requests.post', return_value=mock_response):
|
||||
success, response = client.send_message("test")
|
||||
|
||||
self.assertFalse(success)
|
||||
error_details = response.get("details", "")
|
||||
|
||||
print(f"\nDEBUG: Error details: {error_details}")
|
||||
self.assertNotIn("SuperSecretToken123", error_details)
|
||||
self.assertIn("[REDACTED]", error_details)
|
||||
|
||||
def test_sanitize_token_from_text(self):
|
||||
"""
|
||||
Test the standalone sanitization function.
|
||||
"""
|
||||
webhook_url = "https://discord.com/api/webhooks/123456789/MySecretToken-Part2"
|
||||
|
||||
# Test 1: Simple URL in text
|
||||
text = f"Failed to send to {webhook_url}"
|
||||
sanitized = sanitize_token_from_text(text, webhook_url)
|
||||
self.assertNotIn("MySecretToken-Part2", sanitized)
|
||||
self.assertIn("[REDACTED]", sanitized)
|
||||
|
||||
# Test 2: Token embedded in other text
|
||||
text = "Some error occurred with token MySecretToken-Part2 processing"
|
||||
sanitized = sanitize_token_from_text(text, webhook_url)
|
||||
self.assertNotIn("MySecretToken-Part2", sanitized)
|
||||
self.assertIn("[REDACTED]", sanitized)
|
||||
|
||||
# Test 3: Multiple occurrences
|
||||
text = f"URL: {webhook_url}, Retry: {webhook_url}"
|
||||
sanitized = sanitize_token_from_text(text, webhook_url)
|
||||
self.assertNotIn("MySecretToken-Part2", sanitized)
|
||||
self.assertEqual(sanitized.count("[REDACTED]"), 2)
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
@@ -6,8 +6,7 @@ import json
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from discordsend_utils.workflow_builder import WorkflowBuilder
|
||||
from discordsend_utils.prompt_extractor import extract_prompts_from_workflow
|
||||
from shared.workflow import WorkflowBuilder, extract_prompts_from_workflow
|
||||
|
||||
class TestWorkflowBuilder(unittest.TestCase):
|
||||
"""Tests for the WorkflowBuilder class."""
|
||||
|
||||
Reference in New Issue
Block a user