feat(bot): implement phase 5 - complete bot features

- Add WebSocket reconnection with exponential backoff (1s-60s, ±10% jitter)
- Add error delivery to notify users when jobs fail
- Create templates cog with /template save/load/list/delete commands
- Create history cog with /history (paginated) and /rerun commands
- Fix BotConfig import in bot.py
- Enable templates and history cogs in bot loader

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
AEmotionStudio
2026-01-20 10:11:53 -08:00
co-authored by Claude Opus 4.5
parent 618b9ae4bb
commit b083702eb9
6 changed files with 597 additions and 33 deletions
+4 -4
View File
@@ -5,7 +5,7 @@ import sys
import asyncio
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 +17,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
@@ -92,8 +92,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",
]
+209
View File
@@ -0,0 +1,209 @@
"""History cog for viewing past generations and rerunning them."""
import discord
from discord import app_commands
from discord.ext import commands
import logging
import json
import random
from typing import List
from ..services.permissions import require_permission, Permissions
from ..database.models import JobStatus
from ..embeds.builders import EmbedBuilder
from ...shared.workflow import WorkflowBuilder
logger = logging.getLogger(__name__)
class HistoryPaginator(discord.ui.View):
"""Paginated view for job history."""
def __init__(self, jobs: List, per_page: int = 5):
super().__init__(timeout=180)
self.jobs = jobs
self.per_page = per_page
self.page = 0
self.max_page = (len(jobs) - 1) // per_page if jobs else 0
self._update_buttons()
def _update_buttons(self):
self.prev_button.disabled = self.page <= 0
self.next_button.disabled = self.page >= self.max_page
def get_embed(self) -> discord.Embed:
embed = discord.Embed(title="Generation History", color=discord.Color.blue())
start = self.page * self.per_page
end = start + self.per_page
page_jobs = self.jobs[start:end]
if not page_jobs:
embed.description = "No generation history found."
return embed
lines = []
for job in page_jobs:
status_emoji = {
JobStatus.COMPLETED.value: "✅",
JobStatus.FAILED.value: "❌",
JobStatus.CANCELLED.value: "🚫",
JobStatus.PENDING.value: "⏳",
JobStatus.RUNNING.value: "🔄",
}.get(job.status, "❓")
prompt_preview = (job.positive_prompt or "No prompt")[:50]
if len(job.positive_prompt or "") > 50:
prompt_preview += "..."
timestamp = (
job.created_at.strftime("%Y-%m-%d %H:%M") if job.created_at else "Unknown"
)
lines.append(
f"{status_emoji} **ID: {job.id}** | {timestamp}\n└ {prompt_preview}"
)
embed.description = "\n\n".join(lines)
embed.set_footer(
text=f"Page {self.page + 1}/{self.max_page + 1} | Use /rerun <id> to regenerate"
)
return embed
@discord.ui.button(label="◀ Previous", style=discord.ButtonStyle.secondary)
async def prev_button(
self, interaction: discord.Interaction, button: discord.ui.Button
):
self.page = max(0, self.page - 1)
self._update_buttons()
await interaction.response.edit_message(embed=self.get_embed(), view=self)
@discord.ui.button(label="Next ▶", style=discord.ButtonStyle.secondary)
async def next_button(
self, interaction: discord.Interaction, button: discord.ui.Button
):
self.page = min(self.max_page, self.page + 1)
self._update_buttons()
await interaction.response.edit_message(embed=self.get_embed(), view=self)
class HistoryCog(commands.Cog):
"""View generation history and rerun past jobs."""
def __init__(self, bot):
self.bot = bot
@app_commands.command(name="history", description="View your generation history")
@app_commands.describe(limit="Number of jobs to show (default: 20, max: 50)")
async def history(self, interaction: discord.Interaction, limit: int = 20):
"""Show paginated generation history."""
await interaction.response.defer(ephemeral=True)
limit = min(max(1, limit), 50) # Clamp between 1 and 50
jobs = await self.bot.repository.list_user_jobs(
user_discord_id=str(interaction.user.id), limit=limit
)
if not jobs:
await interaction.followup.send(
"You have no generation history.", ephemeral=True
)
return
view = HistoryPaginator(jobs)
await interaction.followup.send(embed=view.get_embed(), view=view, ephemeral=True)
@app_commands.command(name="rerun", description="Rerun a previous generation")
@app_commands.describe(job_id="The ID of the job to rerun")
@require_permission(Permissions.USER.value)
async def rerun(self, interaction: discord.Interaction, job_id: int):
"""Rerun a previous job with the same parameters."""
await interaction.response.defer()
# Get the original job
original_job = await self.bot.repository.get_job_by_id(job_id)
if not original_job:
await interaction.followup.send(
f"Job ID {job_id} not found.", ephemeral=True
)
return
# Check ownership
if str(original_job.user.discord_id) != str(interaction.user.id):
await interaction.followup.send(
"You can only rerun your own jobs.", ephemeral=True
)
return
# Check if job has required data
if not original_job.workflow_json:
await interaction.followup.send(
"This job cannot be rerun (workflow data not saved).", ephemeral=True
)
return
try:
workflow = json.loads(original_job.workflow_json)
except json.JSONDecodeError:
await interaction.followup.send(
"Failed to parse original workflow.", ephemeral=True
)
return
# Parse parameters
parameters = {}
if original_job.parameters:
try:
parameters = json.loads(original_job.parameters)
except json.JSONDecodeError:
pass
# Generate new seed for rerun
new_seed = random.randint(1, 1000000000000000)
parameters["seed"] = new_seed
# Update workflow with new seed
builder = WorkflowBuilder(workflow)
builder.set_seed(new_seed)
final_workflow = builder.get_workflow()
# Determine delivery and context
server_id = str(interaction.guild_id) if interaction.guild else None
channel_id = str(interaction.channel_id)
try:
# Create new job
job = await self.bot.job_manager.create_job(
user_discord_id=str(interaction.user.id),
workflow=final_workflow,
positive_prompt=original_job.positive_prompt or "",
negative_prompt=original_job.negative_prompt or "",
parameters=parameters,
server_discord_id=server_id,
channel_id=channel_id,
delivery_type=original_job.delivery_type or "channel",
)
embed = EmbedBuilder.job_queued(job)
embed.set_footer(text=f"Rerun of job #{job_id} | New job ID: {job.id}")
await interaction.followup.send(embed=embed)
# Store message ID
original_message = await interaction.original_response()
await self.bot.repository.update_job_message(
job.prompt_id, str(original_message.id)
)
except Exception as e:
logger.error(f"Failed to rerun job: {e}")
await interaction.followup.send(
f"Failed to rerun job: {str(e)}", ephemeral=True
)
async def setup(bot):
await bot.add_cog(HistoryCog(bot))
+223
View File
@@ -0,0 +1,223 @@
"""Template management cog for saving and loading prompt presets."""
import discord
from discord import app_commands
from discord.ext import commands
import logging
from typing import List
from ..services.permissions import require_permission, Permissions
logger = logging.getLogger(__name__)
class TemplateCog(commands.Cog):
"""Manage prompt templates."""
def __init__(self, bot):
self.bot = bot
template_group = app_commands.Group(
name="template", description="Manage prompt templates"
)
@template_group.command(name="save", description="Save a prompt as a template")
@app_commands.describe(
name="Template name (unique per user/server)",
prompt="The positive prompt to save",
negative_prompt="Negative prompt (optional)",
shared="Share with entire server (default: private)",
)
@require_permission(Permissions.GENERATOR.value)
async def template_save(
self,
interaction: discord.Interaction,
name: str,
prompt: str,
negative_prompt: str = "",
shared: bool = False,
):
"""Save current prompt as a named template."""
await interaction.response.defer(ephemeral=True)
# Validate name
if len(name) > 100:
await interaction.followup.send(
"Template name must be 100 characters or less.", ephemeral=True
)
return
# Ensure user exists
await self.bot.repository.get_or_create_user(
str(interaction.user.id), interaction.user.display_name
)
server_id = str(interaction.guild_id) if interaction.guild and shared else None
if server_id:
await self.bot.repository.get_or_create_server(
server_id, interaction.guild.name
)
try:
await self.bot.repository.create_template(
user_discord_id=str(interaction.user.id),
name=name,
positive_prompt=prompt,
negative_prompt=negative_prompt,
server_discord_id=server_id,
)
scope = "server" if shared else "private"
await interaction.followup.send(
f"Saved template **{name}** ({scope}).", ephemeral=True
)
except Exception as e:
if "UNIQUE constraint" in str(e):
await interaction.followup.send(
f"A template named **{name}** already exists. "
"Delete it first or use a different name.",
ephemeral=True,
)
else:
logger.error(f"Failed to save template: {e}")
await interaction.followup.send(
"Failed to save template.", ephemeral=True
)
@template_group.command(name="load", description="Load a saved template")
@app_commands.describe(name="Template name to load")
async def template_load(self, interaction: discord.Interaction, name: str):
"""Load a template and show its contents."""
await interaction.response.defer(ephemeral=True)
server_id = str(interaction.guild_id) if interaction.guild else None
# Try user's private template first
template = await self.bot.repository.get_template(
user_discord_id=str(interaction.user.id), name=name
)
# Try shared server template if not found
if not template and server_id:
templates = await self.bot.repository.list_templates(
user_discord_id=str(interaction.user.id),
server_discord_id=server_id,
include_shared=True,
)
template = next((t for t in templates if t.name == name), None)
if not template:
await interaction.followup.send(
f"Template **{name}** not found.", ephemeral=True
)
return
embed = discord.Embed(title=f"Template: {template.name}", color=discord.Color.blue())
embed.add_field(
name="Prompt", value=template.positive_prompt[:1024], inline=False
)
if template.negative_prompt:
embed.add_field(
name="Negative Prompt",
value=template.negative_prompt[:1024],
inline=False,
)
scope = "Shared" if template.server_id else "Private"
embed.set_footer(text=f"{scope} template | Use /generate with this prompt")
await interaction.followup.send(embed=embed, ephemeral=True)
@template_group.command(name="list", description="List your saved templates")
async def template_list(self, interaction: discord.Interaction):
"""List all available templates."""
await interaction.response.defer(ephemeral=True)
server_id = str(interaction.guild_id) if interaction.guild else None
templates = await self.bot.repository.list_templates(
user_discord_id=str(interaction.user.id),
server_discord_id=server_id,
include_shared=True,
)
if not templates:
await interaction.followup.send(
"You have no saved templates.", ephemeral=True
)
return
embed = discord.Embed(title="Your Templates", color=discord.Color.blue())
private_templates = [t for t in templates if t.server_id is None]
shared_templates = [t for t in templates if t.server_id is not None]
if private_templates:
names = "\n".join([f"• {t.name}" for t in private_templates[:10]])
if len(private_templates) > 10:
names += f"\n...and {len(private_templates) - 10} more"
embed.add_field(name="Private Templates", value=names, inline=False)
if shared_templates:
names = "\n".join([f"• {t.name}" for t in shared_templates[:10]])
if len(shared_templates) > 10:
names += f"\n...and {len(shared_templates) - 10} more"
embed.add_field(name="Server Templates", value=names, inline=False)
await interaction.followup.send(embed=embed, ephemeral=True)
@template_group.command(name="delete", description="Delete a saved template")
@app_commands.describe(name="Template name to delete")
@require_permission(Permissions.GENERATOR.value)
async def template_delete(self, interaction: discord.Interaction, name: str):
"""Delete a template."""
await interaction.response.defer(ephemeral=True)
# Try deleting private template
deleted = await self.bot.repository.delete_template(
user_discord_id=str(interaction.user.id), name=name
)
# Try deleting shared template if private not found
if not deleted and interaction.guild:
deleted = await self.bot.repository.delete_template(
user_discord_id=str(interaction.user.id),
name=name,
server_discord_id=str(interaction.guild_id),
)
if deleted:
await interaction.followup.send(
f"Deleted template **{name}**.", ephemeral=True
)
else:
await interaction.followup.send(
f"Template **{name}** not found or you don't have permission to delete it.",
ephemeral=True,
)
@template_load.autocomplete("name")
@template_delete.autocomplete("name")
async def template_name_autocomplete(
self, interaction: discord.Interaction, current: str
) -> List[app_commands.Choice[str]]:
"""Autocomplete for template names."""
server_id = str(interaction.guild_id) if interaction.guild else None
templates = await self.bot.repository.list_templates(
user_discord_id=str(interaction.user.id),
server_discord_id=server_id,
include_shared=True,
)
# Filter by current input
filtered = [t for t in templates if current.lower() in t.name.lower()]
return [
app_commands.Choice(name=t.name, value=t.name)
for t in filtered[:25] # Discord limit
]
async def setup(bot):
await bot.add_cog(TemplateCog(bot))
+86 -6
View File
@@ -2,10 +2,17 @@ import aiohttp
import logging
import json
import asyncio
import random
from typing import Callable, Coroutine, Any, Dict, List, Optional
logger = logging.getLogger(__name__)
# Reconnection constants
INITIAL_BACKOFF = 1.0 # Initial delay in seconds
MAX_BACKOFF = 60.0 # Maximum delay cap
BACKOFF_MULTIPLIER = 2 # Exponential multiplier
JITTER_FACTOR = 0.1 # +/- 10% randomization
class ComfyUIWebSocket:
"""WebSocket Client for real-time ComfyUI events."""
@@ -27,14 +34,21 @@ class ComfyUIWebSocket:
self._running = False
self._listen_task: Optional[asyncio.Task] = None
# Reconnection state
self._reconnect_attempts = 0
self._should_reconnect = True
self._reconnect_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()
try:
self.ws = await self.session.ws_connect(self.ws_url)
self._running = True
self._should_reconnect = 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}")
except Exception as e:
@@ -45,7 +59,17 @@ class ComfyUIWebSocket:
async def disconnect(self):
"""Disconnect from WebSocket."""
self._should_reconnect = False # Prevent reconnection loop
self._running = False
# Cancel reconnection task if running
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 +113,65 @@ 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:
await self.connect()
logger.info(
f"Successfully reconnected after "
f"{self._reconnect_attempts} attempt(s)."
)
return
except Exception as e:
logger.warning(f"Reconnection attempt failed: {e}")
logger.info("Reconnection loop ended (should_reconnect=False).")
+67 -17
View File
@@ -2,12 +2,14 @@ import discord
import io
import json
import logging
from typing import List, Dict, Any, Union
from typing import Any, Optional, Union
from ..comfyui.client import ComfyUIClient
from ..embeds.builders import EmbedBuilder
logger = logging.getLogger(__name__)
class DeliveryService:
"""Handles delivery of results to Discord."""
@@ -15,7 +17,69 @@ class DeliveryService:
self.bot = bot
self.client = comfy_client
async def deliver_job(self, job: Any): # job: Job model
async def _get_destination(
self, job: Any
) -> Optional[Union[discord.User, discord.TextChannel]]:
"""Get the destination channel or user for a job."""
if job.delivery_type == "dm":
try:
user = self.bot.get_user(int(job.user.discord_id))
if not user:
user = await self.bot.fetch_user(int(job.user.discord_id))
return user
except Exception as e:
logger.error(f"Failed to fetch user for DM: {e}")
return None
if job.channel_id:
destination = self.bot.get_channel(int(job.channel_id))
if destination:
return destination
# Fallback to DM if channel not found
try:
user = self.bot.get_user(int(job.user.discord_id))
if not user:
user = await self.bot.fetch_user(int(job.user.discord_id))
return user
except Exception:
pass
return None
async def deliver_error(self, job: Any, error_message: str) -> bool:
"""
Deliver error notification for a failed job.
Args:
job: The failed Job model instance
error_message: The error message to display
Returns:
True if delivery succeeded, False otherwise
"""
# Truncate error message if too long (Discord embed field limit)
if len(error_message) > 1000:
error_message = error_message[:997] + "..."
embed = EmbedBuilder.job_failed(job, error_message)
destination = await self._get_destination(job)
if destination:
try:
await destination.send(embed=embed)
logger.info(f"Delivered error for job {job.id} to {destination}")
return True
except Exception as e:
logger.error(f"Failed to deliver error: {e}")
return False
else:
logger.error(
f"Could not determine destination for error delivery (job {job.id})"
)
return False
async def deliver_job(self, job: Any): # job: Job model
"""Deliver results for a completed job."""
if not job.output_images:
logger.warning(f"Job {job.id} completed but has no images.")
@@ -43,21 +107,7 @@ class DeliveryService:
logger.warning("No files to upload.")
return
# Determine destination
destination = None
if job.delivery_type == "dm":
user = self.bot.get_user(int(job.user.discord_id)) or await self.bot.fetch_user(int(job.user.discord_id))
destination = user
elif job.channel_id:
destination = self.bot.get_channel(int(job.channel_id))
if not destination:
# Fallback to DM if channel not found?
try:
user = self.bot.get_user(int(job.user.discord_id)) or await self.bot.fetch_user(int(job.user.discord_id))
destination = user
except:
pass
destination = await self._get_destination(job)
if destination:
content = f"Generation complete for <@{job.user.discord_id}>!\n**Prompt:** {job.positive_prompt}"
+8 -6
View File
@@ -196,11 +196,13 @@ class JobManager:
prompt_id = msg.get("prompt_id")
exception_type = msg.get("exception_type", "Unknown Error")
exception_message = msg.get("exception_message", "")
if prompt_id:
error_msg = f"{exception_type}: {exception_message}"
job = await self.repo.update_job_status(prompt_id, JobStatus.FAILED.value, error_message=error_msg)
# Notify user of failure via delivery service?
# Ideally yes, but DeliveryService currently only sends images.
# We might want to expand DeliveryService to handle errors too, or reuse the channel_id to post the failure embed.
pass
job = await self.repo.update_job_status(
prompt_id, JobStatus.FAILED.value, error_message=error_msg
)
# Deliver error notification to user
if job:
await self.delivery.deliver_error(job, error_msg)