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:
co-authored by
Claude Opus 4.5
parent
618b9ae4bb
commit
b083702eb9
+4
-4
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user