Merge pull request #36 from AEmotionStudio/refactor/separation-of-concerns

Refactor/separation of concerns
This commit is contained in:
Æmotion Studio
2026-01-20 16:25:54 -08:00
committed by GitHub
15 changed files with 1426 additions and 56 deletions
+12
View File
@@ -1,3 +1,15 @@
# Sensitive files - NEVER commit these
.env
.env.*
!.env.example
config.yaml
config.yml
*.pem
*.key
secrets.*
credentials.*
# Python
__pycache__/
*.pyc
Errors.md
+37
View File
@@ -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
+6 -7
View File
@@ -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",
]
+1 -1
View File
@@ -7,7 +7,7 @@ from pathlib import Path
from ..embeds.builders import EmbedBuilder
from ..services.permissions import require_permission, Permissions
from ...shared.workflow import WorkflowBuilder
from shared.workflow import WorkflowBuilder
logger = logging.getLogger(__name__)
+209
View File
@@ -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))
+223
View File
@@ -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
View File
@@ -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).")
+67 -17
View File
@@ -2,12 +2,14 @@ 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."""
@@ -15,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.")
@@ -43,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}"
+8 -6
View File
@@ -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)
-3
View File
@@ -2,16 +2,13 @@
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
-2
View File
@@ -11,7 +11,6 @@ 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
@@ -19,7 +18,6 @@ from pathlib import Path
import sys
import datetime
import subprocess
import itertools
import functools
import server
+117
View File
@@ -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()
+216
View File
@@ -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()
+261
View File
@@ -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()
+141
View File
@@ -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()