Merge pull request #36 from AEmotionStudio/refactor/separation-of-concerns
Refactor/separation of concerns
This commit is contained in:
+12
@@ -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
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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).")
|
||||
|
||||
+67
-17
@@ -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}"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,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,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()
|
||||
Reference in New Issue
Block a user