diff --git a/.gitignore b/.gitignore index 69ec3d7..ab613dd 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/CHANGELOG.md b/CHANGELOG.md index 68bebd2..aded04d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.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 diff --git a/bot/bot.py b/bot/bot.py index d5e6f0a..ee2d04f 100644 --- a/bot/bot.py +++ b/bot/bot.py @@ -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", ] diff --git a/bot/cogs/generate.py b/bot/cogs/generate.py index f33be89..e027cb9 100644 --- a/bot/cogs/generate.py +++ b/bot/cogs/generate.py @@ -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__) diff --git a/bot/cogs/history.py b/bot/cogs/history.py new file mode 100644 index 0000000..1fc441d --- /dev/null +++ b/bot/cogs/history.py @@ -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 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)) diff --git a/bot/cogs/templates.py b/bot/cogs/templates.py new file mode 100644 index 0000000..751bb57 --- /dev/null +++ b/bot/cogs/templates.py @@ -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)) diff --git a/bot/comfyui/websocket.py b/bot/comfyui/websocket.py index 612eb12..4247ecc 100644 --- a/bot/comfyui/websocket.py +++ b/bot/comfyui/websocket.py @@ -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).") diff --git a/bot/services/delivery.py b/bot/services/delivery.py index caff2e1..c579770 100644 --- a/bot/services/delivery.py +++ b/bot/services/delivery.py @@ -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}" diff --git a/bot/services/job_manager.py b/bot/services/job_manager.py index 352cbdb..a1b1305 100644 --- a/bot/services/job_manager.py +++ b/bot/services/job_manager.py @@ -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) diff --git a/nodes/image_node.py b/nodes/image_node.py index 01580bb..0ffabc5 100644 --- a/nodes/image_node.py +++ b/nodes/image_node.py @@ -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 diff --git a/nodes/video_node.py b/nodes/video_node.py index cc9ab70..952520d 100644 --- a/nodes/video_node.py +++ b/nodes/video_node.py @@ -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 diff --git a/tests/test_filename_utils.py b/tests/test_filename_utils.py new file mode 100644 index 0000000..520ea8f --- /dev/null +++ b/tests/test_filename_utils.py @@ -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() diff --git a/tests/test_format_utils.py b/tests/test_format_utils.py new file mode 100644 index 0000000..766776f --- /dev/null +++ b/tests/test_format_utils.py @@ -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() diff --git a/tests/test_message_builder.py b/tests/test_message_builder.py new file mode 100644 index 0000000..8614e5f --- /dev/null +++ b/tests/test_message_builder.py @@ -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() diff --git a/tests/test_path_utils.py b/tests/test_path_utils.py new file mode 100644 index 0000000..59db4cf --- /dev/null +++ b/tests/test_path_utils.py @@ -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()