diff --git a/ANALYSIS.md b/ANALYSIS.md new file mode 100644 index 0000000..634a603 --- /dev/null +++ b/ANALYSIS.md @@ -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* diff --git a/PRD.md b/PRD.md new file mode 100644 index 0000000..27483e7 --- /dev/null +++ b/PRD.md @@ -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 diff --git a/bot/__init__.py b/bot/__init__.py new file mode 100644 index 0000000..c54d3b5 --- /dev/null +++ b/bot/__init__.py @@ -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" diff --git a/bot/__main__.py b/bot/__main__.py new file mode 100644 index 0000000..d193287 --- /dev/null +++ b/bot/__main__.py @@ -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() diff --git a/bot/bot.py b/bot/bot.py new file mode 100644 index 0000000..a9ebc5f --- /dev/null +++ b/bot/bot.py @@ -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) diff --git a/bot/cogs/__init__.py b/bot/cogs/__init__.py new file mode 100644 index 0000000..fb076ef --- /dev/null +++ b/bot/cogs/__init__.py @@ -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", +] diff --git a/bot/cogs/generate.py b/bot/cogs/generate.py new file mode 100644 index 0000000..a969e2f --- /dev/null +++ b/bot/cogs/generate.py @@ -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)) diff --git a/bot/comfyui/__init__.py b/bot/comfyui/__init__.py new file mode 100644 index 0000000..e3acbdd --- /dev/null +++ b/bot/comfyui/__init__.py @@ -0,0 +1,6 @@ +"""ComfyUI integration package.""" + +from .client import ComfyUIClient +from .websocket import ComfyUIWebSocket + +__all__ = ["ComfyUIClient", "ComfyUIWebSocket"] diff --git a/bot/comfyui/client.py b/bot/comfyui/client.py new file mode 100644 index 0000000..524a140 --- /dev/null +++ b/bot/comfyui/client.py @@ -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() diff --git a/bot/comfyui/websocket.py b/bot/comfyui/websocket.py new file mode 100644 index 0000000..dd5dc63 --- /dev/null +++ b/bot/comfyui/websocket.py @@ -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.") diff --git a/bot/config.py b/bot/config.py new file mode 100644 index 0000000..637295c --- /dev/null +++ b/bot/config.py @@ -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 diff --git a/bot/data/default_workflow_api.json b/bot/data/default_workflow_api.json new file mode 100644 index 0000000..17c112f --- /dev/null +++ b/bot/data/default_workflow_api.json @@ -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" + } + } +} \ No newline at end of file diff --git a/bot/database/__init__.py b/bot/database/__init__.py new file mode 100644 index 0000000..f585d55 --- /dev/null +++ b/bot/database/__init__.py @@ -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", +] diff --git a/bot/database/models.py b/bot/database/models.py new file mode 100644 index 0000000..428eb64 --- /dev/null +++ b/bot/database/models.py @@ -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"" + + +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"" + + +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"" + + +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"" + + +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"" + + +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"" diff --git a/bot/database/repository.py b/bot/database/repository.py new file mode 100644 index 0000000..cf06669 --- /dev/null +++ b/bot/database/repository.py @@ -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() diff --git a/bot/embeds/__init__.py b/bot/embeds/__init__.py new file mode 100644 index 0000000..46a4ba3 --- /dev/null +++ b/bot/embeds/__init__.py @@ -0,0 +1,5 @@ +"""Discord embed builders.""" + +from .builders import EmbedBuilder + +__all__ = ["EmbedBuilder"] diff --git a/bot/embeds/builders.py b/bot/embeds/builders.py new file mode 100644 index 0000000..911ac95 --- /dev/null +++ b/bot/embeds/builders.py @@ -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 diff --git a/bot/services/__init__.py b/bot/services/__init__.py new file mode 100644 index 0000000..136e5e7 --- /dev/null +++ b/bot/services/__init__.py @@ -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", +] diff --git a/bot/services/delivery.py b/bot/services/delivery.py new file mode 100644 index 0000000..2cf4ac6 --- /dev/null +++ b/bot/services/delivery.py @@ -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}") diff --git a/bot/services/job_manager.py b/bot/services/job_manager.py new file mode 100644 index 0000000..a555389 --- /dev/null +++ b/bot/services/job_manager.py @@ -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 diff --git a/requirements.txt b/requirements.txt index 1a020ff..c040606 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,7 @@ -requests>=2.25.0 \ No newline at end of file +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 \ No newline at end of file diff --git a/tests/bot/__init__.py b/tests/bot/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/bot/test_imports.py b/tests/bot/test_imports.py new file mode 100644 index 0000000..9d4fad5 --- /dev/null +++ b/tests/bot/test_imports.py @@ -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) diff --git a/utils/workflow_builder.py b/utils/workflow_builder.py new file mode 100644 index 0000000..70e153c --- /dev/null +++ b/utils/workflow_builder.py @@ -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 +