feat: Implement Phase 1 & 2 of ComfyUI Companion Bot

This commit is contained in:
AEmotionStudio
2026-01-12 21:21:16 -08:00
parent 29a48ae943
commit 5755166e41
24 changed files with 3002 additions and 1 deletions
+361
View File
@@ -0,0 +1,361 @@
# ComfyUI-DiscordSend Codebase Analysis
## Executive Summary
ComfyUI-DiscordSend is a well-structured custom node extension (~3,824 lines) that bridges ComfyUI's AI image/video generation with Discord via webhooks. The codebase has been recently modularized (v1.1.0) with clear separation of concerns, comprehensive security measures, and solid testing foundations.
---
## Current Architecture Overview
### Project Structure
```
comfyui-discordsend/
├── __init__.py # Node registration (30 lines)
├── discord_image_node.py # Image node (986 lines)
├── discord_video_node.py # Video node (1,549 lines)
├── utils/
│ ├── __init__.py # Public API exports
│ ├── sanitizer.py # Security sanitization (269 lines)
│ ├── discord_api.py # Webhook client (409 lines)
│ ├── github_integration.py # CDN URL archival (137 lines)
│ ├── prompt_extractor.py # Workflow extraction (240 lines)
│ └── logging_config.py # Structured logging (43 lines)
├── tests/
│ └── test_utils.py # Unit tests (142 lines)
└── requirements.txt # Minimal: requests>=2.25.0
```
### Core Capabilities
| Feature | Image Node | Video Node |
|---------|------------|------------|
| Discord Webhook Integration | Yes | Yes |
| Batch Processing | Up to 9 images | Single video |
| Format Options | PNG, JPEG, WebP | GIF, MP4, WebM, ProRes |
| Metadata Embedding | PNG chunks | N/A |
| Workflow Export | Yes | Yes |
| GitHub CDN Archival | Yes | Yes |
| Prompt Extraction | Yes | Yes |
### Strengths
1. **Modular Design**: Clean separation between nodes, API layer, and utilities
2. **Security-First**: Comprehensive sanitization of webhooks and tokens
3. **Minimal Dependencies**: Only `requests>=2.25.0` required
4. **Graceful Degradation**: FFmpeg optional with fallbacks
5. **ComfyUI Integration**: Proper INPUT_TYPES, UI previews, hidden inputs
6. **Test Coverage**: Foundation tests for security-critical code
### Areas for Improvement
1. **Node File Complexity**: Both node files are large (986 and 1,549 lines)
2. **Duplication**: Shared logic between image/video nodes could be abstracted
3. **Test Coverage**: Only utils tested; nodes lack unit tests
4. **Error Recovery**: Some edge cases could use more graceful handling
5. **Configuration**: Hardcoded limits (25MB, 9 images) could be configurable
---
## Refactoring Recommendations
### Priority 1: Extract Shared Node Logic
Both nodes share significant patterns that could be consolidated:
```python
# Proposed: utils/node_base.py
class DiscordNodeBase:
"""Shared functionality for Discord nodes."""
def validate_webhook(self, url): ...
def prepare_discord_message(self, prompts, metadata): ...
def handle_github_integration(self, cdn_urls, config): ...
def sanitize_workflow(self, workflow): ...
```
**Benefits**: Reduced duplication, easier maintenance, consistent behavior
### Priority 2: Configuration Management
```python
# Proposed: utils/config.py
class DiscordConfig:
MAX_FILE_SIZE = 25 * 1024 * 1024 # 25MB
MAX_IMAGES_PER_BATCH = 9
MAX_MESSAGE_LENGTH = 2000
WEBHOOK_TIMEOUT = 30
# ...configurable via environment or config file
```
### Priority 3: Expand Test Coverage
- Add integration tests for node execution
- Mock Discord API for end-to-end testing
- Test video encoding paths
### Priority 4: Type Hints
The codebase lacks comprehensive type hints. Adding them would improve:
- IDE support and autocomplete
- Static analysis with mypy
- Documentation clarity
---
## Companion Discord Bot Analysis
### Feasibility Assessment: **Highly Feasible**
The current architecture actually makes a companion bot quite natural to implement:
| Factor | Assessment | Notes |
|--------|------------|-------|
| **API Abstraction** | Ready | `discord_api.py` already handles Discord communication |
| **Sanitization** | Ready | Security layer is mature and reusable |
| **Dependencies** | Minimal | Would add `discord.py` or `pycord` |
| **Architecture** | Compatible | Modular design allows bot to share utils |
### What a Companion Bot Could Offer
#### 1. **Interactive Queue Management**
```
User: /queue status
Bot: 📊 Your ComfyUI Queue:
• Position 3 of 7
• Estimated time: ~4 minutes
• Current workflow: "portrait_generation_v2"
User: /queue cancel 5
Bot: ✅ Cancelled job #5 (landscape_batch)
```
Currently, users send images after generation. A bot could provide real-time queue visibility and control directly in Discord.
#### 2. **Workflow Triggers from Discord**
```
User: /generate portrait --prompt "cyberpunk warrior" --steps 30
Bot: 🎨 Queued! Job #42
Workflow: portrait_template
Estimated: 2 minutes
[2 minutes later]
Bot: ✨ Job #42 Complete! [4 images attached]
```
**Benefits**:
- No need to open ComfyUI for simple generations
- Mobile-friendly generation triggers
- Preset workflows accessible via slash commands
#### 3. **Prompt Management & Templates**
```
User: /prompt save "hero-shot" "cinematic lighting, dramatic pose, 8k"
Bot: 💾 Saved prompt template "hero-shot"
User: /prompt list
Bot: Your templates:
• hero-shot: "cinematic lighting..."
• anime-style: "anime, cel shaded..."
• photorealistic: "RAW photo, 8k..."
User: /generate using hero-shot --subject "robot warrior"
```
#### 4. **Gallery & History**
```
User: /gallery today
Bot: 📸 Today's Generations (23 images)
[Thumbnail grid with navigation buttons]
User: /history #42
Bot: Job #42 Details:
• Workflow: portrait_v2
• Seed: 12345
• Steps: 30
• [Re-run] [Variations] [Upscale]
```
#### 5. **User Preference Storage**
```
User: /settings default-steps 25
Bot: ✅ Default steps set to 25
User: /settings show
Bot: Your Settings:
• Default steps: 25
• Default sampler: euler_ancestral
• Auto-send to #ai-art: enabled
• Quality preset: high
```
#### 6. **Batch Operations**
```
User: /batch upscale --channel #raw-outputs --count 10
Bot: 🔄 Queued 10 images for upscaling
Progress: ████████░░ 8/10
```
#### 7. **Server Administration**
```
Admin: /discordsend config set-channel #ai-art
Bot: ✅ Default output channel set to #ai-art
Admin: /discordsend stats
Bot: 📊 Server Statistics (This Month):
• Total generations: 1,247
• Top user: @alice (342)
• Peak hour: 8-9 PM
• Avg generation time: 45s
```
### Architecture for Bot Integration
```
┌─────────────────────────────────────────────────────────────┐
│ Discord Server │
│ ┌──────────┐ ┌──────────┐ ┌──────────────────────────┐ │
│ │ Users │ │ Channels │ │ Slash Commands │ │
│ └────┬─────┘ └────┬─────┘ └────────────┬─────────────┘ │
└───────┼─────────────┼────────────────────┼──────────────────┘
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────────────────┐
│ Companion Discord Bot │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Command Handler (slash commands, messages) │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ Queue Manager │ Template Store │ User Prefs │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ ComfyUI API Client (REST/WebSocket) │ │
│ └─────────────────────────────────────────────────────┘ │
└────────────────────────┬────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────┐
│ ComfyUI Server │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ ComfyUI Core + API │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ comfyui-discordsend (existing nodes) │ │
│ │ • DiscordSendSaveImage │ │
│ │ • DiscordSendSaveVideo │ │
│ │ • Shared utils (sanitizer, discord_api, etc.) │ │
│ └─────────────────────────────────────────────────────┘ │
└─────────────────────────────────────────────────────────────┘
```
### Shared Code Strategy
```python
# The bot could import and reuse existing utilities:
from comfyui_discordsend.utils import (
sanitize_json_for_export, # Security
validate_webhook_url, # Validation
DiscordWebhookClient, # API (for fallback)
extract_prompts_from_workflow, # Workflow parsing
)
# New bot-specific modules:
comfyui_discordsend_bot/
├── bot.py # Main bot entry point
├── cogs/
│ ├── generation.py # /generate, /queue commands
│ ├── gallery.py # /gallery, /history commands
│ ├── templates.py # /prompt, /workflow commands
│ └── admin.py # Server configuration
├── comfyui_client.py # ComfyUI API integration
├── database/
│ ├── models.py # User prefs, templates, history
│ └── migrations/
└── config.py # Bot configuration
```
### Benefits Summary
| Benefit | Impact | Complexity |
|---------|--------|------------|
| Mobile generation access | High | Medium |
| Queue visibility | High | Low |
| Prompt templates | Medium | Low |
| Generation history | Medium | Medium |
| Server statistics | Low | Low |
| Batch operations | High | Medium |
| User preferences | Medium | Low |
| Workflow triggers | Very High | High |
### Recommended Bot Tech Stack
| Component | Recommendation | Rationale |
|-----------|----------------|-----------|
| Framework | `discord.py` or `pycord` | Mature, async, slash command support |
| Database | SQLite (small), PostgreSQL (large) | User prefs, history, templates |
| ComfyUI Integration | REST API + WebSocket | Queue status, generation triggers |
| Hosting | Self-hosted alongside ComfyUI | Shared resources, low latency |
---
## Implementation Roadmap
### Phase 1: Foundation (Refactoring)
1. Extract shared node logic into base class
2. Create configuration management module
3. Expand test coverage
4. Add type hints to public APIs
### Phase 2: Bot MVP
1. Set up Discord bot skeleton with slash commands
2. Implement `/generate` with basic workflow trigger
3. Add `/queue` status commands
4. Create simple SQLite storage for user preferences
### Phase 3: Enhanced Features
1. Prompt template system
2. Generation history and gallery
3. Batch operations
4. Server administration commands
### Phase 4: Polish
1. Comprehensive documentation
2. Docker deployment options
3. Configuration UI (web dashboard?)
4. Community workflow sharing
---
## Risk Assessment
| Risk | Likelihood | Mitigation |
|------|------------|------------|
| ComfyUI API changes | Medium | Abstract API layer, version pinning |
| Discord API rate limits | Low | Implement proper rate limiting |
| Security exposure | Medium | Reuse existing sanitization, audit bot commands |
| Maintenance burden | Medium | Clear separation between node and bot repos |
| User adoption | Unknown | Start with high-value features, gather feedback |
---
## Conclusion
Adding a companion Discord bot is **highly feasible** and would significantly enhance the value proposition of ComfyUI-DiscordSend. The current modular architecture provides an excellent foundation, with reusable utilities for security, API communication, and workflow handling.
**Key Recommendation**: Start with a minimal bot that solves one pain point well (e.g., queue visibility or simple generation triggers), then expand based on user feedback. The existing open-source model supports this incremental approach.
The refactoring suggestions above would prepare the codebase for bot integration while also improving the standalone node quality. Consider tackling Phase 1 refactoring before or in parallel with bot development.
---
*Analysis generated: 2026-01-12*
*Codebase version: 1.1.0*
+409
View File
@@ -0,0 +1,409 @@
# Product Requirements Document: ComfyUI-DiscordSend Companion Bot
## Overview
A Discord bot that enables users to trigger ComfyUI image generations, manage queues, save prompt templates, and browse generation history—all from within Discord.
**Project**: `comfyui-discordsend-bot`
**Parent Extension**: `comfyui-discordsend` (v1.1.0)
**License**: GPL-3.0 (matching parent)
**Repository**: Same repo (`bot/` folder)
---
## Problem Statement
Currently, users must:
1. Open ComfyUI's web interface to queue generations
2. Wait for completion with no mobile-friendly status updates
3. Manually manage prompts across sessions
4. Rely on webhooks for one-way delivery (no interaction)
The companion bot solves these by bringing full generation control into Discord.
---
## Target Users
| User Type | Use Case |
|-----------|----------|
| **Individual Creators** | Run personal ComfyUI instance, trigger generations from phone/Discord |
| **Community Servers** | Shared ComfyUI where multiple users queue generations |
---
## MVP Features
### 1. Generation Triggers (`/generate`)
```
/generate prompt:"cyberpunk cityscape" [negative:"blurry"] [template:my-style]
[steps:30] [cfg:7.5] [seed:12345] [delivery:dm]
```
- Submit generation requests directly from Discord
- Support positive/negative prompts
- Optional parameter overrides (steps, CFG, seed, dimensions)
- Load from saved templates
- Choose delivery: channel or DM
### 2. Queue Management (`/queue`)
```
/queue view # See current queue with positions
/queue status 42 # Detailed status of job #42
/queue cancel 42 # Cancel your job (or any with admin)
/queue clear # Clear your pending jobs
```
- Real-time queue position updates
- Progress bars during generation
- Cancel pending/running jobs
- Per-user queue limits (configurable)
### 3. Prompt Templates (`/template`)
```
/template save name:"portrait-style" prompt:"cinematic lighting, 8k"
/template list
/template load name:"portrait-style"
/template delete name:"portrait-style"
```
- Save reusable prompt presets
- Private templates (user-only)
- Shared templates (server-wide, optional)
- Include negative prompts and parameters
### 4. Generation History (`/history`)
```
/history list [limit:10]
/history view 42
/history rerun 42 [delivery:dm]
```
- Browse past generations
- View prompts and parameters used
- Re-run with same or modified settings
- Filter by status (completed, failed)
### 5. DM Delivery
- Send completed images directly to user's DMs
- Per-request flag: `delivery:dm` or `delivery:channel`
- User preference for default delivery method
---
## Permission System
Role-based access control using Discord server roles:
| Level | Capabilities |
|-------|-------------|
| **user** | View queue, view own history |
| **generator** | Generate, templates, cancel own jobs |
| **admin** | All + manage roles, cancel any job, configure bot |
Configuration via `/admin setroles generator:@Role admin:@Role`
---
## Technical Architecture
### Integration Model
```
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
│ Discord Users │────►│ Companion Bot │────►│ ComfyUI │
│ (slash cmds) │◄────│ (discord.py) │◄────│ (localhost) │
└─────────────────┘ └────────┬────────┘ └─────────────────┘
│
┌────────▼────────┐
│ SQLite │
│ (jobs, users, │
│ templates) │
└─────────────────┘
```
### Deployment Model
- **Same machine** as ComfyUI (localhost API)
- Shares Python environment with ComfyUI
- Reuses existing `utils/` modules from comfyui-discordsend
### ComfyUI Integration
| Method | Purpose |
|--------|---------|
| REST API | Submit workflows, get history, fetch images |
| WebSocket | Real-time progress updates, completion events |
Key endpoints: `/prompt`, `/queue`, `/history`, `/view`, `/interrupt`
### Data Storage (SQLite)
| Table | Purpose |
|-------|---------|
| `users` | Discord user preferences |
| `servers` | Guild configuration |
| `server_roles` | Permission mappings |
| `templates` | Saved prompt presets |
| `jobs` | Generation history and queue |
| `workflows` | Default workflow storage |
---
## Project Structure
```
comfyui-discordsend/
├── utils/ # Existing shared utilities (reuse)
├── bot/ # NEW: Discord bot package
│ ├── __main__.py # Entry point
│ ├── bot.py # Main bot class
│ ├── config.py # Configuration
│ ├── database/
│ │ ├── models.py # SQLAlchemy models
│ │ └── repository.py # Data access
│ ├── comfyui/
│ │ ├── client.py # REST client
│ │ └── websocket.py # WebSocket client
│ ├── cogs/
│ │ ├── generate.py # /generate
│ │ ├── queue.py # /queue
│ │ ├── templates.py # /template
│ │ ├── history.py # /history
│ │ └── admin.py # /admin
│ ├── services/
│ │ ├── job_manager.py # Job lifecycle
│ │ ├── delivery.py # Result delivery
│ │ └── permissions.py # RBAC
│ └── embeds/
│ └── builders.py # Discord embeds
└── requirements.txt # Bot dependencies
```
---
## Code Sharing Strategy
Reuse from existing `utils/`:
| Module | Bot Usage |
|--------|-----------|
| `sanitizer.py` | Sanitize workflows before storage/display |
| `prompt_extractor.py` | Extract prompts for history/template views |
| `discord_api.py` | Retry patterns, rate limit handling |
| `logging_config.py` | Consistent logging |
New shared utility:
- `utils/workflow_builder.py` - Modify workflow JSON (set prompts, seed, steps, etc.)
---
## Dependencies
```
# Bot-specific
discord.py>=2.3.0
aiohttp>=3.9.0
sqlalchemy>=2.0.0
aiosqlite>=0.19.0
pydantic>=2.0.0
pyyaml>=6.0.0
# Shared with extension
requests>=2.25.0
```
---
## Configuration
Environment variables (`.env`):
```
DISCORDBOT_DISCORD_TOKEN=your_bot_token
DISCORDBOT_COMFYUI_URL=http://127.0.0.1:8188
```
Bot config (`bot/config.yaml`):
```yaml
defaults:
max_queue_per_user: 3
progress_update_interval: 2.0
workflow_path: workflows/default.json
```
---
## User Experience Flow
### Generation Flow
1. User runs `/generate prompt:"a robot"`
2. Bot validates permissions, checks queue limits
3. Bot submits workflow to ComfyUI
4. Bot posts "Queued" embed with position
5. WebSocket updates trigger progress bar edits
6. On completion, bot fetches images from ComfyUI
7. Bot delivers images to channel or DM (user's choice)
### Error Handling
| Scenario | Response |
|----------|----------|
| Permission denied | Ephemeral message explaining required role |
| Queue full | Ephemeral message with current limit |
| ComfyUI offline | Ephemeral message suggesting to check server |
| Generation failed | Update embed with error, log details |
---
## Implementation Phases
### Phase 1: Foundation
- Project structure and configuration
- Database schema and migrations
- ComfyUI REST client
- Basic bot lifecycle
### Phase 2: Core Generation
- `/generate` command
- WebSocket integration for progress
- Job manager and tracking
- Result delivery (channel/DM)
### Phase 3: Queue & Permissions
- `/queue` commands
- Role-based permission system
- `/admin` configuration commands
- Per-user queue limits
### Phase 4: Templates & History
- `/template` CRUD commands
- `/history` browsing and re-run
- Autocomplete for template names
### Phase 5: Polish
- Comprehensive error handling
- Progress embeds with previews
- Unit and integration tests
- Documentation and README
---
## Verification Plan
1. **Unit Tests**: Config, permissions, workflow builder
2. **Integration Tests**: Database ops, ComfyUI client (mocked)
3. **Manual E2E Testing**:
- Generate image via `/generate`
- Verify progress updates in Discord
- Confirm delivery to channel and DM
- Test queue cancellation
- Test template save/load cycle
- Test history re-run
---
## Files to Create
| File | Purpose |
|------|---------|
| `bot/__init__.py` | Package init with path setup |
| `bot/__main__.py` | Entry point (`python -m bot`) |
| `bot/bot.py` | Main bot class |
| `bot/config.py` | Configuration management |
| `bot/database/models.py` | SQLAlchemy models |
| `bot/database/repository.py` | Data access layer |
| `bot/comfyui/client.py` | REST API client |
| `bot/comfyui/websocket.py` | WebSocket client |
| `bot/cogs/generate.py` | /generate command |
| `bot/cogs/queue.py` | /queue commands |
| `bot/cogs/templates.py` | /template commands |
| `bot/cogs/history.py` | /history commands |
| `bot/cogs/admin.py` | /admin commands |
| `bot/services/job_manager.py` | Job lifecycle management |
| `bot/services/delivery.py` | Result delivery |
| `bot/services/permissions.py` | RBAC |
| `bot/embeds/builders.py` | Discord embed builders |
| `utils/workflow_builder.py` | Shared workflow manipulation |
---
## Database Schema
```sql
-- Users table
CREATE TABLE users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
discord_id TEXT UNIQUE NOT NULL,
username TEXT NOT NULL,
default_delivery TEXT DEFAULT 'channel',
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
-- Servers (guilds) table
CREATE TABLE servers (
id INTEGER PRIMARY KEY AUTOINCREMENT,
discord_id TEXT UNIQUE NOT NULL,
name TEXT NOT NULL,
default_channel_id TEXT,
max_queue_per_user INTEGER DEFAULT 3,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
-- Server roles for permissions
CREATE TABLE server_roles (
id INTEGER PRIMARY KEY AUTOINCREMENT,
server_id INTEGER NOT NULL,
role_discord_id TEXT NOT NULL,
permission_level TEXT NOT NULL,
FOREIGN KEY (server_id) REFERENCES servers(id)
);
-- Templates (user prompt presets)
CREATE TABLE templates (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
server_id INTEGER,
name TEXT NOT NULL,
positive_prompt TEXT NOT NULL,
negative_prompt TEXT DEFAULT '',
parameters TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id)
);
-- Jobs (generation history)
CREATE TABLE jobs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
prompt_id TEXT UNIQUE NOT NULL,
user_id INTEGER NOT NULL,
server_id INTEGER,
channel_id TEXT,
status TEXT DEFAULT 'pending',
positive_prompt TEXT,
negative_prompt TEXT,
parameters TEXT,
output_images TEXT,
error_message TEXT,
delivery_type TEXT,
message_id TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
completed_at TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id)
);
```
---
## Success Metrics
- Users can generate images without opening ComfyUI web UI
- Queue visibility reduces "is it working?" uncertainty
- Templates reduce repetitive prompt typing
- History enables easy iteration on generations
- DM delivery provides private results in shared servers
+17
View File
@@ -0,0 +1,17 @@
"""
ComfyUI-DiscordSend Companion Bot
A Discord bot for triggering ComfyUI generations, managing queues,
saving prompt templates, and browsing generation history.
"""
import sys
from pathlib import Path
# Add parent directory to path for importing shared utils
_parent_dir = Path(__file__).parent.parent
if str(_parent_dir) not in sys.path:
sys.path.insert(0, str(_parent_dir))
__version__ = "0.1.0"
__author__ = "AEmotionStudio"
+42
View File
@@ -0,0 +1,42 @@
import asyncio
import logging
import sys
import os
from pathlib import Path
# Add project root to python path to allow imports
project_root = Path(__file__).resolve().parent.parent
sys.path.append(str(project_root))
from bot.config import Config
from bot.bot import ComfyUIBot
from utils.logging_config import setup_logging
def main():
# Setup logging
setup_logging()
logger = logging.getLogger("bot")
# Load configuration
try:
config = Config()
except Exception as e:
logger.critical(f"Failed to load configuration: {e}")
return
if not config.discord_token:
logger.critical("Discord token not found! Set DISCORDBOT_DISCORD_TOKEN env var or config.yaml")
return
# Initialize and run bot
bot = ComfyUIBot(config)
try:
bot.run(config.discord_token)
except KeyboardInterrupt:
logger.info("Bot stopped by user.")
except Exception as e:
logger.critical(f"Bot crashed: {e}")
if __name__ == "__main__":
main()
+122
View File
@@ -0,0 +1,122 @@
import discord
from discord.ext import commands
import logging
import sys
import asyncio
from pathlib import Path
from .config import Config
from .database.repository import Repository
from .comfyui.client import ComfyUIClient
from .comfyui.websocket import ComfyUIWebSocket
logger = logging.getLogger(__name__)
class ComfyUIBot(commands.Bot):
"""
Main Bot Class for ComfyUI Companion.
"""
def __init__(self, config: Config):
intents = discord.Intents.default()
intents.message_content = True # Needed for some commands if not pure slash
intents.members = True # Useful for permission checks
super().__init__(
command_prefix=commands.when_mentioned,
intents=intents,
help_command=None,
description="ComfyUI Companion Bot"
)
self.config = config
# Database
self.repository = Repository(config.database_url)
# ComfyUI Clients
self.comfy_client = ComfyUIClient(base_url=config.comfyui_url)
self.comfy_ws = ComfyUIWebSocket(base_url=config.comfyui_url)
# Services
from .services.delivery import DeliveryService
from .services.job_manager import JobManager
self.delivery_service = DeliveryService(self, self.comfy_client)
self.job_manager = JobManager(
self.repository,
self.comfy_client,
self.comfy_ws,
self.delivery_service
)
async def setup_hook(self):
"""Async setup before bot starts."""
logger.info(f"Setting up bot for {self.user}...")
# Initialize Database
await self.repository.init_db()
logger.info("Database initialized.")
# Connect to ComfyUI WebSocket
try:
await self.comfy_ws.connect()
await self.job_manager.start()
except Exception as e:
logger.warning(f"Could not connect to ComfyUI WebSocket at start: {e}")
# We don't crash, as ComfyUI might come up later
# Load Cogs
await self._load_cogs()
# Sync Commands
# In production, sync might be manual or per-guild to avoid rate limits
# For simplicity in this self-hosted bot, we'll sync global
try:
synced = await self.tree.sync()
logger.info(f"Synced {len(synced)} slash commands.")
except Exception as e:
logger.error(f"Failed to sync commands: {e}")
async def _load_cogs(self):
"""Load extensions from cogs directory."""
# We need to use dotted path: bot.cogs.generate
cogs_dir = Path(__file__).parent / "cogs"
# Extensions to load
extensions = [
"bot.cogs.generate",
# "bot.cogs.queue",
# "bot.cogs.templates",
# "bot.cogs.history",
# "bot.cogs.admin",
]
for ext in extensions:
try:
# Check if file exists first to avoid confusing errors if we haven't created it yet
# (Since we are building incrementally)
module_name = ext.split(".")[-1]
if (cogs_dir / f"{module_name}.py").exists():
await self.load_extension(ext)
logger.info(f"Loaded extension: {ext}")
else:
logger.debug(f"Skipping extension {ext} (file not found)")
except Exception as e:
logger.error(f"Failed to load extension {ext}: {e}")
async def close(self):
"""Cleanup on shutdown."""
logger.info("Shutting down bot...")
await self.comfy_client.close()
await self.comfy_ws.disconnect()
await self.repository.close()
await super().close()
async def on_ready(self):
logger.info(f"Bot logged in as {self.user} (ID: {self.user.id})")
logger.info(f"Connected to {len(self.guilds)} guilds")
async def on_command_error(self, ctx, error):
"""Global error handler for prefix commands (if any)."""
logger.error(f"Command error: {error}", exc_info=False)
+15
View File
@@ -0,0 +1,15 @@
"""Discord bot cogs (command groups)."""
from .generate import GenerateCog
from .queue import QueueCog
from .templates import TemplateCog
from .history import HistoryCog
from .admin import AdminCog
__all__ = [
"GenerateCog",
"QueueCog",
"TemplateCog",
"HistoryCog",
"AdminCog",
]
+118
View File
@@ -0,0 +1,118 @@
import discord
from discord import app_commands
from discord.ext import commands
import logging
import json
from pathlib import Path
from ..embeds.builders import EmbedBuilder
from ...utils.workflow_builder import WorkflowBuilder
logger = logging.getLogger(__name__)
class GenerateCog(commands.Cog):
def __init__(self, bot):
self.bot = bot
@app_commands.command(name="generate", description="Generate an image using ComfyUI")
@app_commands.describe(
prompt="The positive prompt to generate",
negative_prompt="Aspects to avoid (optional)",
seed="Seed for generation (optional)",
steps="Number of steps (optional)",
cfg="CFG Scale (optional)",
delivery="Delivery method (channel or dm)"
)
@app_commands.choices(delivery=[
app_commands.Choice(name="Current Channel", value="channel"),
app_commands.Choice(name="Direct Message", value="dm")
])
async def generate(self, interaction: discord.Interaction,
prompt: str,
negative_prompt: str = "",
seed: int = None,
steps: int = None,
cfg: float = None,
delivery: app_commands.Choice[str] = None):
await interaction.response.defer()
# Load default workflow
# TODO: Move this path to config or database
workflow_path = Path(__file__).parent.parent / "data" / "default_workflow_api.json"
if not workflow_path.exists():
await interaction.followup.send("❌ Error: Default workflow not found.", ephemeral=True)
return
try:
with open(workflow_path, "r") as f:
workflow_json = json.load(f)
except Exception as e:
logger.error(f"Failed to load workflow: {e}")
await interaction.followup.send("❌ Error: Failed to load workflow configuration.", ephemeral=True)
return
# Prepare parameters for tracking
parameters = {
"seed": seed,
"steps": steps,
"cfg": cfg
}
# Modify workflow
builder = WorkflowBuilder(workflow_json)
builder.set_prompt(prompt, negative_prompt)
if seed is not None:
builder.set_seed(seed)
else:
# Random seed if not provided
import random
generated_seed = random.randint(1, 1000000000000000)
builder.set_seed(generated_seed)
parameters["seed"] = generated_seed # Track actual seed
if steps is not None:
builder.set_steps(steps)
if cfg is not None:
builder.set_cfg(cfg)
final_workflow = builder.get_workflow()
# Determine delivery method
delivery_method = delivery.value if delivery else "channel"
# Check server context
server_id = str(interaction.guild_id) if interaction.guild else None
channel_id = str(interaction.channel_id)
try:
# Create Job
job = await self.bot.job_manager.create_job(
user_discord_id=str(interaction.user.id),
workflow=final_workflow,
positive_prompt=prompt,
negative_prompt=negative_prompt,
parameters=parameters,
server_discord_id=server_id,
channel_id=channel_id,
delivery_type=delivery_method
)
# Send Queued Embed
embed = EmbedBuilder.job_queued(job)
await interaction.followup.send(embed=embed)
# Store the interaction message ID if we want to update it later
# (JobManager could use this to update the specific message)
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 start generation: {e}")
await interaction.followup.send(f"❌ Error starting generation: {str(e)}", ephemeral=True)
async def setup(bot):
await bot.add_cog(GenerateCog(bot))
+6
View File
@@ -0,0 +1,6 @@
"""ComfyUI integration package."""
from .client import ComfyUIClient
from .websocket import ComfyUIWebSocket
__all__ = ["ComfyUIClient", "ComfyUIWebSocket"]
+107
View File
@@ -0,0 +1,107 @@
import aiohttp
import logging
from typing import Optional, Dict, List, Any
import json
import uuid
logger = logging.getLogger(__name__)
class ComfyUIClient:
"""Async Client for interacting with ComfyUI REST API."""
def __init__(self, base_url: str = "http://127.0.0.1:8188"):
self.base_url = base_url.rstrip("/")
self.session: Optional[aiohttp.ClientSession] = None
async def _get_session(self) -> aiohttp.ClientSession:
if self.session is None or self.session.closed:
self.session = aiohttp.ClientSession()
return self.session
async def close(self):
if self.session and not self.session.closed:
await self.session.close()
async def check_status(self) -> bool:
"""Check if ComfyUI server is reachable."""
try:
session = await self._get_session()
async with session.get(f"{self.base_url}/system_stats") as response:
return response.status == 200
except Exception as e:
logger.warning(f"Failed to connect to ComfyUI: {e}")
return False
async def get_system_stats(self) -> Dict[str, Any]:
"""Get system statistics."""
session = await self._get_session()
async with session.get(f"{self.base_url}/system_stats") as response:
response.raise_for_status()
return await response.json()
async def get_queue(self) -> Dict[str, Any]:
"""Get current queue status."""
session = await self._get_session()
async with session.get(f"{self.base_url}/queue") as response:
response.raise_for_status()
return await response.json()
async def get_history(self, prompt_id: str) -> Dict[str, Any]:
"""Get history for a specific prompt ID."""
session = await self._get_session()
async with session.get(f"{self.base_url}/history/{prompt_id}") as response:
response.raise_for_status()
return await response.json()
async def queue_prompt(self, workflow: Dict[str, Any], client_id: str) -> Dict[str, Any]:
"""
Queue a workflow for generation.
Args:
workflow: The workflow JSON object (API format)
client_id: Unique client ID for WebSocket correlation
"""
session = await self._get_session()
payload = {
"prompt": workflow,
"client_id": client_id
}
async with session.post(f"{self.base_url}/prompt", json=payload) as response:
response.raise_for_status()
return await response.json()
async def interrupt(self):
"""Interrupt currently executing prompt."""
session = await self._get_session()
async with session.post(f"{self.base_url}/interrupt") as response:
try:
response.raise_for_status()
except Exception as e:
logger.error(f"Failed to interrupt: {e}")
async def delete_queue_item(self, prompt_id: str):
"""Remove an item from queue."""
session = await self._get_session()
payload = {"delete": [prompt_id]}
async with session.post(f"{self.base_url}/queue", json=payload) as response:
response.raise_for_status()
async def get_image(self, filename: str, subfolder: str = "", type: str = "output") -> bytes:
"""Fetch a generated image."""
session = await self._get_session()
params = {"filename": filename, "subfolder": subfolder, "type": type}
async with session.get(f"{self.base_url}/view", params=params) as response:
response.raise_for_status()
return await response.read()
async def upload_image(self, image_data: bytes, filename: str, subfolder: str = ""):
"""Upload an image to ComfyUI (input folder)."""
session = await self._get_session()
data = aiohttp.FormData()
data.add_field("image", image_data, filename=filename)
if subfolder:
data.add_field("subfolder", subfolder)
async with session.post(f"{self.base_url}/upload/image", data=data) as response:
response.raise_for_status()
return await response.json()
+94
View File
@@ -0,0 +1,94 @@
import aiohttp
import logging
import json
import asyncio
from typing import Callable, Coroutine, Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class ComfyUIWebSocket:
"""WebSocket Client for real-time ComfyUI events."""
def __init__(self, base_url: str = "ws://127.0.0.1:8188", client_id: str = ""):
# Convert http/https to ws/wss if needed
if base_url.startswith("http://"):
base_url = base_url.replace("http://", "ws://")
elif base_url.startswith("https://"):
base_url = base_url.replace("https://", "wss://")
self.ws_url = f"{base_url.rstrip('/')}/ws"
if client_id:
self.ws_url += f"?clientId={client_id}"
self.ws: Optional[aiohttp.ClientWebSocketResponse] = None
self.session: Optional[aiohttp.ClientSession] = None
self._callbacks: Dict[str, List[Callable[[Dict[str, Any]], Coroutine[Any, Any, None]]]] = {}
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()
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}")
raise
async def disconnect(self):
"""Disconnect from WebSocket."""
self._running = False
if self.ws:
await self.ws.close()
if self.session:
await self.session.close()
if self._listen_task:
try:
await self._listen_task
except asyncio.CancelledError:
pass
def add_listener(self, event_type: str, callback: Callable[[Dict[str, Any]], Coroutine[Any, Any, None]]):
"""Register a callback for an event type."""
if event_type not in self._callbacks:
self._callbacks[event_type] = []
self._callbacks[event_type].append(callback)
async def _listen(self):
"""Listen loop for incoming messages."""
if not self.ws:
return
try:
async for msg in self.ws:
if msg.type == aiohttp.WSMsgType.TEXT:
try:
data = json.loads(msg.data)
event_type = data.get("type", "unknown")
# Some messages pack content in 'data', others at top level
# ComfyUI typically sends {type: "event_name", data: {...}, sid: "..."}
handlers = self._callbacks.get(event_type, [])
for handler in handlers:
try:
await handler(data)
except Exception as e:
logger.error(f"Error in WebSocket handler for {event_type}: {e}")
except json.JSONDecodeError:
logger.warning(f"Received invalid JSON: {msg.data}")
elif msg.type == aiohttp.WSMsgType.ERROR:
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
finally:
if self._running:
logger.info("WebSocket listener stopped unexpectedly.")
+212
View File
@@ -0,0 +1,212 @@
"""
Configuration management for the Discord bot.
Configuration is loaded from (in priority order):
1. Environment variables (highest priority)
2. Config file (bot/config.yaml)
3. Defaults (lowest priority)
"""
import os
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
import yaml
@dataclass
class DiscordConfig:
"""Discord-related configuration."""
token: str = ""
application_id: Optional[str] = None
@dataclass
class ComfyUIConfig:
"""ComfyUI connection configuration."""
url: str = "http://127.0.0.1:8188"
ws_url: str = "ws://127.0.0.1:8188/ws"
timeout: int = 30
@dataclass
class DefaultsConfig:
"""Default values for bot operations."""
max_queue_per_user: int = 3
progress_update_interval: float = 2.0
workflow_path: Optional[str] = None
default_steps: int = 20
default_cfg: float = 7.0
default_width: int = 512
default_height: int = 512
@dataclass
class DatabaseConfig:
"""Database configuration."""
url: str = ""
def __post_init__(self):
if not self.url:
# Default to SQLite in bot/data directory
bot_dir = Path(__file__).parent
data_dir = bot_dir / "data"
data_dir.mkdir(exist_ok=True)
self.url = f"sqlite+aiosqlite:///{data_dir}/bot.db"
@dataclass
class SecurityConfig:
"""Security configuration."""
allowed_guilds: list[int] = field(default_factory=list)
@dataclass
class BotConfig:
"""Main bot configuration container."""
discord: DiscordConfig = field(default_factory=DiscordConfig)
comfyui: ComfyUIConfig = field(default_factory=ComfyUIConfig)
defaults: DefaultsConfig = field(default_factory=DefaultsConfig)
database: DatabaseConfig = field(default_factory=DatabaseConfig)
security: SecurityConfig = field(default_factory=SecurityConfig)
@classmethod
def load(cls, config_path: Optional[Path] = None) -> "BotConfig":
"""
Load configuration from file and environment variables.
Args:
config_path: Optional path to config.yaml file
Returns:
Loaded BotConfig instance
"""
config = cls()
# Load from config file if exists
if config_path is None:
config_path = Path(__file__).parent / "config.yaml"
if config_path.exists():
config._load_from_file(config_path)
# Override with environment variables
config._load_from_env()
return config
def _load_from_file(self, path: Path) -> None:
"""Load configuration from YAML file."""
with open(path) as f:
data = yaml.safe_load(f) or {}
# Discord config
if "discord" in data:
discord_data = data["discord"]
if "token" in discord_data:
self.discord.token = discord_data["token"]
if "application_id" in discord_data:
self.discord.application_id = discord_data["application_id"]
# ComfyUI config
if "comfyui" in data:
comfyui_data = data["comfyui"]
if "url" in comfyui_data:
self.comfyui.url = comfyui_data["url"]
if "ws_url" in comfyui_data:
self.comfyui.ws_url = comfyui_data["ws_url"]
if "timeout" in comfyui_data:
self.comfyui.timeout = comfyui_data["timeout"]
# Defaults config
if "defaults" in data:
defaults_data = data["defaults"]
if "max_queue_per_user" in defaults_data:
self.defaults.max_queue_per_user = defaults_data["max_queue_per_user"]
if "progress_update_interval" in defaults_data:
self.defaults.progress_update_interval = defaults_data["progress_update_interval"]
if "workflow_path" in defaults_data:
self.defaults.workflow_path = defaults_data["workflow_path"]
if "default_steps" in defaults_data:
self.defaults.default_steps = defaults_data["default_steps"]
if "default_cfg" in defaults_data:
self.defaults.default_cfg = defaults_data["default_cfg"]
if "default_width" in defaults_data:
self.defaults.default_width = defaults_data["default_width"]
if "default_height" in defaults_data:
self.defaults.default_height = defaults_data["default_height"]
# Database config
if "database" in data:
db_data = data["database"]
if "url" in db_data:
self.database.url = db_data["url"]
# Security config
if "security" in data:
security_data = data["security"]
if "allowed_guilds" in security_data:
self.security.allowed_guilds = security_data["allowed_guilds"] or []
def _load_from_env(self) -> None:
"""Load configuration from environment variables."""
# Discord
if token := os.getenv("DISCORDBOT_DISCORD_TOKEN"):
self.discord.token = token
if app_id := os.getenv("DISCORDBOT_APPLICATION_ID"):
self.discord.application_id = app_id
# ComfyUI
if url := os.getenv("DISCORDBOT_COMFYUI_URL"):
self.comfyui.url = url
if ws_url := os.getenv("DISCORDBOT_COMFYUI_WS_URL"):
self.comfyui.ws_url = ws_url
if timeout := os.getenv("DISCORDBOT_COMFYUI_TIMEOUT"):
self.comfyui.timeout = int(timeout)
# Database
if db_url := os.getenv("DISCORDBOT_DATABASE_URL"):
self.database.url = db_url
# Defaults
if max_queue := os.getenv("DISCORDBOT_MAX_QUEUE_PER_USER"):
self.defaults.max_queue_per_user = int(max_queue)
if workflow := os.getenv("DISCORDBOT_WORKFLOW_PATH"):
self.defaults.workflow_path = workflow
def validate(self) -> list[str]:
"""
Validate the configuration.
Returns:
List of validation error messages (empty if valid)
"""
errors = []
if not self.discord.token:
errors.append("Discord token is required (set DISCORDBOT_DISCORD_TOKEN)")
if not self.comfyui.url:
errors.append("ComfyUI URL is required")
return errors
# Global config instance (lazy loaded)
_config: Optional[BotConfig] = None
def get_config() -> BotConfig:
"""Get the global configuration instance."""
global _config
if _config is None:
_config = BotConfig.load()
return _config
def reload_config(config_path: Optional[Path] = None) -> BotConfig:
"""Reload configuration from disk."""
global _config
_config = BotConfig.load(config_path)
return _config
+107
View File
@@ -0,0 +1,107 @@
{
"3": {
"inputs": {
"seed": 156680208700286,
"steps": 20,
"cfg": 8,
"sampler_name": "euler",
"scheduler": "normal",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"6",
0
],
"negative": [
"7",
0
],
"latent_image": [
"5",
0
]
},
"class_type": "KSampler",
"_meta": {
"title": "KSampler"
}
},
"4": {
"inputs": {
"ckpt_name": "v1-5-pruned-emaonly.ckpt"
},
"class_type": "CheckpointLoaderSimple",
"_meta": {
"title": "Load Checkpoint"
}
},
"5": {
"inputs": {
"width": 512,
"height": 512,
"batch_size": 1
},
"class_type": "EmptyLatentImage",
"_meta": {
"title": "Empty Latent Image"
}
},
"6": {
"inputs": {
"text": "beautiful scenery nature glass bottle landscape, , purple galaxy bottle,",
"clip": [
"4",
1
]
},
"class_type": "CLIPTextEncode",
"_meta": {
"title": "Positive Prompt"
}
},
"7": {
"inputs": {
"text": "text, watermark",
"clip": [
"4",
1
]
},
"class_type": "CLIPTextEncode",
"_meta": {
"title": "Negative Prompt"
}
},
"8": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"4",
2
]
},
"class_type": "VAEDecode",
"_meta": {
"title": "VAE Decode"
}
},
"9": {
"inputs": {
"filename_prefix": "ComfyUI",
"images": [
"8",
0
]
},
"class_type": "SaveImage",
"_meta": {
"title": "Save Image"
}
}
}
+15
View File
@@ -0,0 +1,15 @@
"""Database package for bot persistence."""
from .models import Base, User, Server, ServerRole, Template, Job, Workflow
from .repository import Repository
__all__ = [
"Base",
"User",
"Server",
"ServerRole",
"Template",
"Job",
"Workflow",
"Repository",
]
+222
View File
@@ -0,0 +1,222 @@
"""
SQLAlchemy ORM models for bot persistence.
Tables:
- users: Discord user preferences
- servers: Guild configuration
- server_roles: Permission role mappings
- templates: Saved prompt presets
- jobs: Generation history and queue
- workflows: Default workflow storage
"""
from datetime import datetime
from enum import Enum
from typing import Optional
from sqlalchemy import (
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.orm import DeclarativeBase, relationship
class Base(DeclarativeBase):
"""Base class for all models."""
pass
class DeliveryType(str, Enum):
"""Delivery method for generated images."""
CHANNEL = "channel"
DM = "dm"
class JobStatus(str, Enum):
"""Status of a generation job."""
PENDING = "pending"
RUNNING = "running"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
class PermissionLevel(str, Enum):
"""Permission levels for role-based access."""
USER = "user"
GENERATOR = "generator"
ADMIN = "admin"
class User(Base):
"""Discord user preferences and data."""
__tablename__ = "users"
id = Column(Integer, primary_key=True, autoincrement=True)
discord_id = Column(String(20), unique=True, nullable=False, index=True)
username = Column(String(100), nullable=False)
default_delivery = Column(String(10), default=DeliveryType.CHANNEL.value)
default_workflow_id = Column(Integer, ForeignKey("workflows.id"), nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# Relationships
templates = relationship("Template", back_populates="user", cascade="all, delete-orphan")
jobs = relationship("Job", back_populates="user", cascade="all, delete-orphan")
default_workflow = relationship("Workflow", foreign_keys=[default_workflow_id])
def __repr__(self) -> str:
return f"<User(id={self.id}, discord_id={self.discord_id}, username={self.username})>"
class Server(Base):
"""Discord server (guild) configuration."""
__tablename__ = "servers"
id = Column(Integer, primary_key=True, autoincrement=True)
discord_id = Column(String(20), unique=True, nullable=False, index=True)
name = Column(String(100), nullable=False)
default_channel_id = Column(String(20), nullable=True)
enabled = Column(Boolean, default=True)
max_queue_per_user = Column(Integer, default=3)
default_workflow_id = Column(Integer, ForeignKey("workflows.id"), nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# Relationships
roles = relationship("ServerRole", back_populates="server", cascade="all, delete-orphan")
templates = relationship("Template", back_populates="server", cascade="all, delete-orphan")
jobs = relationship("Job", back_populates="server")
default_workflow = relationship("Workflow", foreign_keys=[default_workflow_id])
def __repr__(self) -> str:
return f"<Server(id={self.id}, discord_id={self.discord_id}, name={self.name})>"
class ServerRole(Base):
"""Role-based permission mapping for servers."""
__tablename__ = "server_roles"
id = Column(Integer, primary_key=True, autoincrement=True)
server_id = Column(Integer, ForeignKey("servers.id"), nullable=False)
role_discord_id = Column(String(20), nullable=False)
permission_level = Column(String(20), nullable=False)
# Relationships
server = relationship("Server", back_populates="roles")
__table_args__ = (
UniqueConstraint("server_id", "role_discord_id", name="uq_server_role"),
Index("idx_server_roles_server", "server_id"),
)
def __repr__(self) -> str:
return f"<ServerRole(server_id={self.server_id}, role={self.role_discord_id}, level={self.permission_level})>"
class Workflow(Base):
"""Stored workflow configurations."""
__tablename__ = "workflows"
id = Column(Integer, primary_key=True, autoincrement=True)
name = Column(String(100), nullable=False)
description = Column(Text, nullable=True)
workflow_json = Column(Text, nullable=False)
is_default = Column(Boolean, default=False)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
def __repr__(self) -> str:
return f"<Workflow(id={self.id}, name={self.name}, is_default={self.is_default})>"
class Template(Base):
"""User-saved prompt templates."""
__tablename__ = "templates"
id = Column(Integer, primary_key=True, autoincrement=True)
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
server_id = Column(Integer, ForeignKey("servers.id"), nullable=True) # NULL = private
name = Column(String(100), nullable=False)
positive_prompt = Column(Text, nullable=False)
negative_prompt = Column(Text, default="")
parameters = Column(Text, nullable=True) # JSON: {steps, cfg, seed, etc.}
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# Relationships
user = relationship("User", back_populates="templates")
server = relationship("Server", back_populates="templates")
__table_args__ = (
Index("idx_templates_user", "user_id"),
Index("idx_templates_server", "server_id"),
UniqueConstraint("user_id", "server_id", "name", name="uq_template_name"),
)
def __repr__(self) -> str:
return f"<Template(id={self.id}, name={self.name}, user_id={self.user_id})>"
class Job(Base):
"""Generation job tracking and history."""
__tablename__ = "jobs"
id = Column(Integer, primary_key=True, autoincrement=True)
prompt_id = Column(String(64), unique=True, nullable=False, index=True)
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
server_id = Column(Integer, ForeignKey("servers.id"), nullable=True)
channel_id = Column(String(20), nullable=True)
# Status tracking
status = Column(String(20), default=JobStatus.PENDING.value, index=True)
queue_position = Column(Integer, nullable=True)
progress = Column(Integer, default=0)
progress_max = Column(Integer, default=0)
# Generation parameters
positive_prompt = Column(Text, nullable=True)
negative_prompt = Column(Text, nullable=True)
parameters = Column(Text, nullable=True) # JSON
workflow_json = Column(Text, nullable=True)
# Results
output_images = Column(Text, nullable=True) # JSON array
error_message = Column(Text, nullable=True)
# Delivery
delivery_type = Column(String(10), default=DeliveryType.CHANNEL.value)
message_id = Column(String(20), nullable=True) # Discord message for updates
# Timestamps
created_at = Column(DateTime, default=datetime.utcnow)
started_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
# Relationships
user = relationship("User", back_populates="jobs")
server = relationship("Server", back_populates="jobs")
__table_args__ = (
Index("idx_jobs_user", "user_id"),
Index("idx_jobs_status", "status"),
Index("idx_jobs_created", "created_at"),
)
@property
def duration(self) -> Optional[float]:
"""Get job duration in seconds."""
if self.started_at and self.completed_at:
return (self.completed_at - self.started_at).total_seconds()
return None
def __repr__(self) -> str:
return f"<Job(id={self.id}, prompt_id={self.prompt_id}, status={self.status})>"
+645
View File
@@ -0,0 +1,645 @@
"""
Data access layer for bot database operations.
Provides async CRUD operations for all database models.
"""
import json
from datetime import datetime
from typing import Optional
from sqlalchemy import select, update, delete, func
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
from sqlalchemy.orm import selectinload
from .models import (
Base,
User,
Server,
ServerRole,
Template,
Job,
Workflow,
JobStatus,
PermissionLevel,
)
class Repository:
"""Async repository for database operations."""
def __init__(self, database_url: str):
"""
Initialize the repository.
Args:
database_url: SQLAlchemy async database URL
"""
self.engine = create_async_engine(database_url, echo=False)
self.async_session = async_sessionmaker(
self.engine,
class_=AsyncSession,
expire_on_commit=False,
)
async def init_db(self) -> None:
"""Create all database tables."""
async with self.engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
async def close(self) -> None:
"""Close the database connection."""
await self.engine.dispose()
# ==================== User Operations ====================
async def get_or_create_user(
self,
discord_id: str,
username: str,
) -> User:
"""Get existing user or create new one."""
async with self.async_session() as session:
result = await session.execute(
select(User).where(User.discord_id == discord_id)
)
user = result.scalar_one_or_none()
if user is None:
user = User(discord_id=discord_id, username=username)
session.add(user)
await session.commit()
await session.refresh(user)
elif user.username != username:
user.username = username
await session.commit()
return user
async def get_user(self, discord_id: str) -> Optional[User]:
"""Get user by Discord ID."""
async with self.async_session() as session:
result = await session.execute(
select(User).where(User.discord_id == discord_id)
)
return result.scalar_one_or_none()
async def update_user_delivery(
self,
discord_id: str,
delivery_type: str,
) -> Optional[User]:
"""Update user's default delivery preference."""
async with self.async_session() as session:
result = await session.execute(
select(User).where(User.discord_id == discord_id)
)
user = result.scalar_one_or_none()
if user:
user.default_delivery = delivery_type
await session.commit()
return user
# ==================== Server Operations ====================
async def get_or_create_server(
self,
discord_id: str,
name: str,
) -> Server:
"""Get existing server or create new one."""
async with self.async_session() as session:
result = await session.execute(
select(Server).where(Server.discord_id == discord_id)
)
server = result.scalar_one_or_none()
if server is None:
server = Server(discord_id=discord_id, name=name)
session.add(server)
await session.commit()
await session.refresh(server)
elif server.name != name:
server.name = name
await session.commit()
return server
async def get_server(self, discord_id: str) -> Optional[Server]:
"""Get server by Discord ID."""
async with self.async_session() as session:
result = await session.execute(
select(Server)
.options(selectinload(Server.roles))
.where(Server.discord_id == discord_id)
)
return result.scalar_one_or_none()
async def update_server_channel(
self,
discord_id: str,
channel_id: str,
) -> Optional[Server]:
"""Update server's default output channel."""
async with self.async_session() as session:
result = await session.execute(
select(Server).where(Server.discord_id == discord_id)
)
server = result.scalar_one_or_none()
if server:
server.default_channel_id = channel_id
await session.commit()
return server
async def update_server_queue_limit(
self,
discord_id: str,
limit: int,
) -> Optional[Server]:
"""Update server's per-user queue limit."""
async with self.async_session() as session:
result = await session.execute(
select(Server).where(Server.discord_id == discord_id)
)
server = result.scalar_one_or_none()
if server:
server.max_queue_per_user = limit
await session.commit()
return server
# ==================== Role Operations ====================
async def set_server_role(
self,
server_discord_id: str,
role_discord_id: str,
permission_level: str,
) -> ServerRole:
"""Set or update a role's permission level for a server."""
async with self.async_session() as session:
# Get server
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if not server:
raise ValueError(f"Server {server_discord_id} not found")
# Check for existing role mapping
role_result = await session.execute(
select(ServerRole).where(
ServerRole.server_id == server.id,
ServerRole.role_discord_id == role_discord_id,
)
)
role = role_result.scalar_one_or_none()
if role:
role.permission_level = permission_level
else:
role = ServerRole(
server_id=server.id,
role_discord_id=role_discord_id,
permission_level=permission_level,
)
session.add(role)
await session.commit()
await session.refresh(role)
return role
async def get_server_roles(self, server_discord_id: str) -> list[ServerRole]:
"""Get all role mappings for a server."""
async with self.async_session() as session:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if not server:
return []
result = await session.execute(
select(ServerRole).where(ServerRole.server_id == server.id)
)
return list(result.scalars().all())
async def delete_server_role(
self,
server_discord_id: str,
role_discord_id: str,
) -> bool:
"""Remove a role mapping."""
async with self.async_session() as session:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if not server:
return False
result = await session.execute(
delete(ServerRole).where(
ServerRole.server_id == server.id,
ServerRole.role_discord_id == role_discord_id,
)
)
await session.commit()
return result.rowcount > 0
# ==================== Template Operations ====================
async def create_template(
self,
user_discord_id: str,
name: str,
positive_prompt: str,
negative_prompt: str = "",
parameters: Optional[dict] = None,
server_discord_id: Optional[str] = None,
) -> Template:
"""Create a new prompt template."""
async with self.async_session() as session:
# Get user
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
raise ValueError(f"User {user_discord_id} not found")
# Get server if provided
server_id = None
if server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
server_id = server.id
template = Template(
user_id=user.id,
server_id=server_id,
name=name,
positive_prompt=positive_prompt,
negative_prompt=negative_prompt,
parameters=json.dumps(parameters) if parameters else None,
)
session.add(template)
await session.commit()
await session.refresh(template)
return template
async def get_template(
self,
user_discord_id: str,
name: str,
server_discord_id: Optional[str] = None,
) -> Optional[Template]:
"""Get a template by name."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
return None
# Build query
query = select(Template).where(
Template.user_id == user.id,
Template.name == name,
)
if server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
query = query.where(Template.server_id == server.id)
else:
query = query.where(Template.server_id.is_(None))
result = await session.execute(query)
return result.scalar_one_or_none()
async def list_templates(
self,
user_discord_id: str,
server_discord_id: Optional[str] = None,
include_shared: bool = True,
) -> list[Template]:
"""List templates for a user."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
return []
# Private templates
query = select(Template).where(
Template.user_id == user.id,
Template.server_id.is_(None),
)
result = await session.execute(query)
templates = list(result.scalars().all())
# Shared templates in server
if include_shared and server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
shared_query = select(Template).where(
Template.server_id == server.id
)
shared_result = await session.execute(shared_query)
templates.extend(shared_result.scalars().all())
return templates
async def delete_template(
self,
user_discord_id: str,
name: str,
server_discord_id: Optional[str] = None,
) -> bool:
"""Delete a template."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
return False
query = delete(Template).where(
Template.user_id == user.id,
Template.name == name,
)
if server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
query = query.where(Template.server_id == server.id)
else:
query = query.where(Template.server_id.is_(None))
result = await session.execute(query)
await session.commit()
return result.rowcount > 0
# ==================== Job Operations ====================
async def create_job(
self,
prompt_id: str,
user_discord_id: str,
positive_prompt: str,
negative_prompt: str = "",
parameters: Optional[dict] = None,
workflow_json: Optional[str] = None,
delivery_type: str = "channel",
server_discord_id: Optional[str] = None,
channel_id: Optional[str] = None,
) -> Job:
"""Create a new generation job."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
raise ValueError(f"User {user_discord_id} not found")
server_id = None
if server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
server_id = server.id
job = Job(
prompt_id=prompt_id,
user_id=user.id,
server_id=server_id,
channel_id=channel_id,
positive_prompt=positive_prompt,
negative_prompt=negative_prompt,
parameters=json.dumps(parameters) if parameters else None,
workflow_json=workflow_json,
delivery_type=delivery_type,
)
session.add(job)
await session.commit()
await session.refresh(job)
return job
async def get_job(self, prompt_id: str) -> Optional[Job]:
"""Get a job by prompt ID."""
async with self.async_session() as session:
result = await session.execute(
select(Job)
.options(selectinload(Job.user))
.where(Job.prompt_id == prompt_id)
)
return result.scalar_one_or_none()
async def get_job_by_id(self, job_id: int) -> Optional[Job]:
"""Get a job by internal ID."""
async with self.async_session() as session:
result = await session.execute(
select(Job)
.options(selectinload(Job.user))
.where(Job.id == job_id)
)
return result.scalar_one_or_none()
async def update_job_status(
self,
prompt_id: str,
status: str,
error_message: Optional[str] = None,
output_images: Optional[list[str]] = None,
) -> Optional[Job]:
"""Update job status."""
async with self.async_session() as session:
result = await session.execute(
select(Job).where(Job.prompt_id == prompt_id)
)
job = result.scalar_one_or_none()
if not job:
return None
job.status = status
if status == JobStatus.RUNNING.value:
job.started_at = datetime.utcnow()
elif status in (JobStatus.COMPLETED.value, JobStatus.FAILED.value, JobStatus.CANCELLED.value):
job.completed_at = datetime.utcnow()
if error_message is not None:
job.error_message = error_message
if output_images is not None:
job.output_images = json.dumps(output_images)
await session.commit()
return job
async def update_job_progress(
self,
prompt_id: str,
progress: int,
progress_max: int,
) -> Optional[Job]:
"""Update job progress."""
async with self.async_session() as session:
result = await session.execute(
select(Job).where(Job.prompt_id == prompt_id)
)
job = result.scalar_one_or_none()
if job:
job.progress = progress
job.progress_max = progress_max
await session.commit()
return job
async def update_job_message(
self,
prompt_id: str,
message_id: str,
) -> Optional[Job]:
"""Update the Discord message ID for a job."""
async with self.async_session() as session:
result = await session.execute(
select(Job).where(Job.prompt_id == prompt_id)
)
job = result.scalar_one_or_none()
if job:
job.message_id = message_id
await session.commit()
return job
async def list_user_jobs(
self,
user_discord_id: str,
limit: int = 10,
status: Optional[str] = None,
) -> list[Job]:
"""List jobs for a user."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
return []
query = (
select(Job)
.where(Job.user_id == user.id)
.order_by(Job.created_at.desc())
.limit(limit)
)
if status:
query = query.where(Job.status == status)
result = await session.execute(query)
return list(result.scalars().all())
async def count_user_pending_jobs(
self,
user_discord_id: str,
server_discord_id: Optional[str] = None,
) -> int:
"""Count pending/running jobs for a user."""
async with self.async_session() as session:
user_result = await session.execute(
select(User).where(User.discord_id == user_discord_id)
)
user = user_result.scalar_one_or_none()
if not user:
return 0
query = (
select(func.count(Job.id))
.where(Job.user_id == user.id)
.where(Job.status.in_([JobStatus.PENDING.value, JobStatus.RUNNING.value]))
)
if server_discord_id:
server_result = await session.execute(
select(Server).where(Server.discord_id == server_discord_id)
)
server = server_result.scalar_one_or_none()
if server:
query = query.where(Job.server_id == server.id)
result = await session.execute(query)
return result.scalar() or 0
async def get_pending_jobs(self) -> list[Job]:
"""Get all pending jobs ordered by creation time."""
async with self.async_session() as session:
result = await session.execute(
select(Job)
.options(selectinload(Job.user))
.where(Job.status.in_([JobStatus.PENDING.value, JobStatus.RUNNING.value]))
.order_by(Job.created_at)
)
return list(result.scalars().all())
# ==================== Workflow Operations ====================
async def save_workflow(
self,
name: str,
workflow_json: str,
description: Optional[str] = None,
is_default: bool = False,
) -> Workflow:
"""Save a workflow configuration."""
async with self.async_session() as session:
# If setting as default, unset other defaults
if is_default:
await session.execute(
update(Workflow).where(Workflow.is_default == True).values(is_default=False)
)
workflow = Workflow(
name=name,
workflow_json=workflow_json,
description=description,
is_default=is_default,
)
session.add(workflow)
await session.commit()
await session.refresh(workflow)
return workflow
async def get_default_workflow(self) -> Optional[Workflow]:
"""Get the default workflow."""
async with self.async_session() as session:
result = await session.execute(
select(Workflow).where(Workflow.is_default == True)
)
return result.scalar_one_or_none()
async def get_workflow(self, name: str) -> Optional[Workflow]:
"""Get a workflow by name."""
async with self.async_session() as session:
result = await session.execute(
select(Workflow).where(Workflow.name == name)
)
return result.scalar_one_or_none()
+5
View File
@@ -0,0 +1,5 @@
"""Discord embed builders."""
from .builders import EmbedBuilder
__all__ = ["EmbedBuilder"]
+72
View File
@@ -0,0 +1,72 @@
import discord
import logging
from typing import Optional, List
from datetime import datetime
class EmbedBuilder:
"""Helper for building Discord embeds."""
@staticmethod
def job_queued(job, position: int = 0) -> discord.Embed:
"""Embed for queued job."""
embed = discord.Embed(
title="🎨 Generation Queued",
description=f"**Prompt:** {job.positive_prompt}",
color=discord.Color.blue(),
timestamp=datetime.utcnow()
)
embed.add_field(name="Queue Position", value=str(position) if position > 0 else "Pending...", inline=True)
embed.add_field(name="Status", value="Waiting to start...", inline=True)
if job.negative_prompt:
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
embed.set_footer(text=f"Job ID: {job.id}")
return embed
@staticmethod
def job_progress(job, progress: int, max_progress: int) -> discord.Embed:
"""Embed for running job with progress."""
percent = int((progress / max_progress) * 100) if max_progress > 0 else 0
bars = "█" * (percent // 10) + "░" * (10 - (percent // 10))
embed = discord.Embed(
title="🎨 Generating...",
description=f"**Prompt:** {job.positive_prompt}",
color=discord.Color.orange(),
timestamp=datetime.utcnow()
)
embed.add_field(name="Progress", value=f"`{bars}` {percent}%", inline=False)
if job.negative_prompt:
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
embed.set_footer(text=f"Job ID: {job.id}")
return embed
@staticmethod
def job_completed(job, image_count: int) -> discord.Embed:
"""Embed for completed job."""
embed = discord.Embed(
title="✨ Generation Complete!",
description=f"**Prompt:** {job.positive_prompt}",
color=discord.Color.green(),
timestamp=datetime.utcnow()
)
embed.add_field(name="Images", value=f"{image_count} generated", inline=True)
embed.add_field(name="Duration", value=f"{job.execution_time:.1f}s" if hasattr(job, 'execution_time') and job.execution_time else "Done", inline=True)
if job.negative_prompt:
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
embed.set_footer(text=f"Job ID: {job.id}")
return embed
@staticmethod
def job_failed(job, error_message: str) -> discord.Embed:
"""Embed for failed job."""
embed = discord.Embed(
title="❌ Generation Failed",
description=f"**Prompt:** {job.positive_prompt}",
color=discord.Color.red(),
timestamp=datetime.utcnow()
)
embed.add_field(name="Error", value=f"```{error_message}```", inline=False)
embed.set_footer(text=f"Job ID: {job.id}")
return embed
+13
View File
@@ -0,0 +1,13 @@
"""Bot services for business logic."""
from .permissions import PermissionService, Permissions, require_permission
from .job_manager import JobManager
from .delivery import DeliveryService
__all__ = [
"PermissionService",
"Permissions",
"require_permission",
"JobManager",
"DeliveryService",
]
+69
View File
@@ -0,0 +1,69 @@
import discord
import io
import logging
from typing import List, Dict, Any, Union
from ..comfyui.client import ComfyUIClient
logger = logging.getLogger(__name__)
class DeliveryService:
"""Handles delivery of results to Discord."""
def __init__(self, bot: discord.Client, comfy_client: ComfyUIClient):
self.bot = bot
self.client = comfy_client
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.")
return
try:
images_data = json.loads(job.output_images)
except json.JSONDecodeError:
logger.error(f"Failed to parse output images for job {job.id}")
return
files = []
for img_meta in images_data:
filename = img_meta.get("filename")
subfolder = img_meta.get("subfolder", "")
img_type = img_meta.get("type", "output")
try:
img_bytes = await self.client.get_image(filename, subfolder, img_type)
files.append(discord.File(io.BytesIO(img_bytes), filename=filename))
except Exception as e:
logger.error(f"Failed to download image {filename}: {e}")
if not files:
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
if destination:
content = f"Generation complete for <@{job.user.discord_id}>!\n**Prompt:** {job.positive_prompt}"
try:
await destination.send(content=content, files=files)
logger.info(f"Delivered job {job.id} to {destination}")
except Exception as e:
logger.error(f"Failed to send message: {e}")
else:
logger.error(f"Could not determine destination for job {job.id}")
+158
View File
@@ -0,0 +1,158 @@
import asyncio
import logging
import json
import uuid
from datetime import datetime
from typing import Optional, Dict, List, Any
from ..database.repository import Repository
from ..database.models import JobStatus, Job
from ..comfyui.client import ComfyUIClient
from ..comfyui.websocket import ComfyUIWebSocket
from .delivery import DeliveryService
logger = logging.getLogger(__name__)
class JobManager:
"""Manages the lifecycle of generation jobs."""
def __init__(self, repository: Repository, comfy_client: ComfyUIClient, comfy_ws: ComfyUIWebSocket, delivery_service: DeliveryService):
self.repo = repository
self.client = comfy_client
self.ws = comfy_ws
self.delivery = delivery_service
self.client_id = str(uuid.uuid4())
# In-memory mapping of prompt_id -> current status buffer
self._active_jobs = {}
async def start(self):
"""Start listening to WebSocket events."""
self.ws.add_listener("status", self._on_status)
self.ws.add_listener("execution_start", self._on_execution_start)
self.ws.add_listener("executing", self._on_executing)
self.ws.add_listener("executed", self._on_executed)
self.ws.add_listener("execution_error", self._on_execution_error)
self.ws.add_listener("progress", self._on_progress)
logger.info(f"JobManager started with client_id: {self.client_id}")
async def create_job(self,
user_discord_id: str,
workflow: Dict[str, Any],
positive_prompt: str,
negative_prompt: str = "",
parameters: Optional[Dict] = None,
server_discord_id: Optional[str] = None,
channel_id: Optional[str] = None,
delivery_type: str = "channel") -> Job:
"""Submit a job to ComfyUI and database."""
# 1. Submit to ComfyUI
response = await self.client.queue_prompt(workflow, self.client_id)
prompt_id = response.get("prompt_id")
if not prompt_id:
raise ValueError("Failed to get prompt_id from ComfyUI")
# 2. Create DB entry
job = await self.repo.create_job(
prompt_id=prompt_id,
user_discord_id=user_discord_id,
server_discord_id=server_discord_id,
channel_id=channel_id,
positive_prompt=positive_prompt,
negative_prompt=negative_prompt,
parameters=parameters,
workflow_json=json.dumps(workflow),
delivery_type=delivery_type
)
logger.info(f"Created job {job.id} (prompt_id: {prompt_id}) for user {user_discord_id}")
return job
async def cancel_job(self, job_id: int) -> bool:
"""Cancel a job."""
job = await self.repo.get_job_by_id(job_id)
if not job:
return False
if job.status in [JobStatus.COMPLETED.value, JobStatus.FAILED.value, JobStatus.CANCELLED.value]:
return False
if job.status == JobStatus.RUNNING.value:
await self.client.interrupt()
# Remove from queue if pending
try:
await self.client.delete_queue_item(job.prompt_id)
except:
pass
await self.repo.update_job_status(job.prompt_id, JobStatus.CANCELLED.value)
return True
# -- WebSocket Event Handlers --
async def _on_status(self, data: Dict[str, Any]):
pass
async def _on_execution_start(self, data: Dict[str, Any]):
prompt_id = data.get("data", {}).get("prompt_id")
if prompt_id:
await self.repo.update_job_status(prompt_id, JobStatus.RUNNING.value)
async def _on_executing(self, data: Dict[str, Any]):
pass
async def _on_progress(self, data: Dict[str, Any]):
msg = data.get("data", {})
prompt_id = msg.get("prompt_id")
value = msg.get("value")
max_val = msg.get("max")
if prompt_id and value is not None and max_val is not None:
await self.repo.update_job_progress(prompt_id, value, max_val)
async def _on_executed(self, data: Dict[str, Any]):
"""
Handle execution completion of a node.
If it contains images, we assume it's a relevant output.
We accumulate images and update job status.
"""
msg = data.get("data", {})
prompt_id = msg.get("prompt_id")
output = msg.get("output", {})
if prompt_id and "images" in output:
# Found images
images = output["images"]
# Update DB with images and mark as completed
# NOTE: In complex workflows, there might be multiple outputs.
# Ideally we check if this is the last one or something.
# But normally 'executed' with images means we got something.
# We'll mark as completed for now. If multiple exist, latest wins.
job = await self.repo.update_job_status(
prompt_id,
JobStatus.COMPLETED.value,
output_images=images
)
if job:
logger.info(f"Job {job.id} completed. Delivering results...")
await self.delivery.deliver_job(job)
async def _on_execution_error(self, data: Dict[str, Any]):
msg = data.get("data", {})
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
+7 -1
View File
@@ -1 +1,7 @@
requests>=2.25.0
requests>=2.25.0
discord.py>=2.3.0
aiohttp>=3.9.0
sqlalchemy>=2.0.0
aiosqlite>=0.19.0
pydantic>=2.0.0
pyyaml>=6.0.0
View File
+39
View File
@@ -0,0 +1,39 @@
import sys
import os
from pathlib import Path
# Add project root to python path
project_root = Path(__file__).resolve().parent.parent.parent
sys.path.append(str(project_root))
print(f"Testing imports from {project_root}")
try:
import bot.config
print("✅ bot.config imported")
import bot.database.models
print("✅ bot.database.models imported")
import bot.database.repository
print("✅ bot.database.repository imported")
import bot.comfyui.client
print("✅ bot.comfyui.client imported")
import bot.services.job_manager
print("✅ bot.services.job_manager imported")
import bot.cogs.generate
print("✅ bot.cogs.generate imported")
import bot.bot
print("✅ bot.bot imported")
print("All modules imported successfully.")
except ImportError as e:
print(f"❌ ImportError: {e}")
sys.exit(1)
except Exception as e:
print(f"❌ Error: {e}")
sys.exit(1)
+147
View File
@@ -0,0 +1,147 @@
import json
import random
import logging
from typing import Dict, Any, Optional, Tuple, List
logger = logging.getLogger(__name__)
class WorkflowBuilder:
"""Helper to manipulate ComfyUI workflow JSONs."""
def __init__(self, workflow_json: Dict[str, Any]):
self.workflow = workflow_json
# Create lookups
self._nodes = self.workflow
if "nodes" in self.workflow and isinstance(self.workflow["nodes"], list):
# Handle "graph" format vs "api" format if needed.
# But usually for API we stick to the {node_id: node_data} format.
# If input is graph format, it might need conversion or distinct handling.
# Assuming API format for now as that's what's sent to /prompt.
pass
@classmethod
def from_json_string(cls, json_str: str) -> 'WorkflowBuilder':
"""Load from JSON int."""
return cls(json.loads(json_str))
def get_workflow(self) -> Dict[str, Any]:
"""Get the current workflow dict."""
return self.workflow
def set_prompt(self, positive: str, negative: Optional[str] = None) -> None:
"""
Attempt to set positive and negative prompts.
Heuristics:
- Look for CLIPTextEncode nodes.
- Often one is connected to KSampler 'positive' and one to 'negative'.
- Or look for custom titles like 'Positive Prompt', 'Negative Prompt'.
"""
# Simple heuristic: Find CLIPTextEncode nodes
# If we have title/coloring, we can use that.
# Otherwise, we might need graph traversal to see what connects to KSampler.
# For MVP, let's assume standard ComfyUI structure or look for specific titles first
positive_node_id = self._find_node_by_title("Positive Prompt")
negative_node_id = self._find_node_by_title("Negative Prompt")
# Fallback: Find KSampler and trace back
if not positive_node_id or not negative_node_id:
ksampler_id, ksampler = self._find_node_by_class("KSampler")
if ksampler:
# KSampler inputs: model, positive, negative, latent_image
if not positive_node_id:
positive_node_id = self._trace_input(ksampler, "positive")
if not negative_node_id:
negative_node_id = self._trace_input(ksampler, "negative")
if positive_node_id:
self._update_node_input(positive_node_id, "text", positive)
else:
logger.warning("Could not identify Positive Prompt node.")
if negative and negative_node_id:
self._update_node_input(negative_node_id, "text", negative)
elif negative:
logger.warning("Could not identify Negative Prompt node.")
def set_seed(self, seed: int) -> int:
"""Set seed on KSampler nodes or Seed nodes."""
# Find KSampler or anything with a 'seed' widget
updated = False
for node_id, node in self.workflow.items():
if "inputs" in node:
if "seed" in node["inputs"]:
# Ensure it's an int widget, not a link
if isinstance(node["inputs"]["seed"], (int, float)) or (isinstance(node["inputs"]["seed"], str) and node["inputs"]["seed"].isdigit()):
node["inputs"]["seed"] = seed
updated = True
if "noise_seed" in node["inputs"]:
# Some nodes call it noise_seed
if isinstance(node["inputs"]["noise_seed"], (int, float)):
node["inputs"]["noise_seed"] = seed
updated = True
if not updated:
logger.warning("Could not find any seed inputs to update.")
return seed
def set_image_dimensions(self, width: int, height: int) -> None:
"""Set width and height on EmptyLatentImage nodes."""
node_id, _ = self._find_node_by_class("EmptyLatentImage")
if node_id:
self._update_node_input(node_id, "width", width)
self._update_node_input(node_id, "height", height)
def set_steps(self, steps: int) -> None:
"""Set steps on KSampler."""
ksampler_ids = self._find_nodes_by_class("KSampler")
for nid in ksampler_ids:
self._update_node_input(nid, "steps", steps)
def set_cfg(self, cfg: float) -> None:
"""Set CFG scale on KSampler."""
ksampler_ids = self._find_nodes_by_class("KSampler")
for nid in ksampler_ids:
self._update_node_input(nid, "cfg", cfg)
def _find_node_by_title(self, title: str) -> Optional[str]:
"""Find node by its custom title (`_meta.title`)."""
for node_id, node in self.workflow.items():
if "_meta" in node and node["_meta"].get("title") == title:
return node_id
return None
def _find_node_by_class(self, class_type: str) -> Tuple[Optional[str], Optional[Dict]]:
"""Find first node of a specific class type."""
for node_id, node in self.workflow.items():
if node.get("class_type") == class_type:
return node_id, node
return None, None
def _find_nodes_by_class(self, class_type: str) -> List[str]:
"""Find all nodes of a specific class type."""
ids = []
for node_id, node in self.workflow.items():
if node.get("class_type") == class_type:
ids.append(node_id)
return ids
def _trace_input(self, node: Dict, input_name: str) -> Optional[str]:
"""
Trace back an input link to find the source node.
Input format in API JSON: "input_name": ["source_node_id", slot_index]
"""
if "inputs" not in node or input_name not in node["inputs"]:
return None
link = node["inputs"][input_name]
# Link structure: [node_id, slot_idx]
if isinstance(link, list) and len(link) == 2:
return str(link[0])
return None
def _update_node_input(self, node_id: str, input_name: str, value: Any) -> None:
if node_id in self.workflow and "inputs" in self.workflow[node_id]:
self.workflow[node_id]["inputs"][input_name] = value