diff --git a/.Jules/palette.md b/.Jules/palette.md new file mode 100644 index 0000000..95ab4b4 --- /dev/null +++ b/.Jules/palette.md @@ -0,0 +1,3 @@ +## 2026-01-10 - Enhanced Tooltip Guidance +**Learning:** Users often miss helpful features like Markdown support in Discord messages or critical warnings because tooltips are generic. +**Action:** Use specific, actionable language in tooltips (e.g., "Supports Discord Markdown") and visual cues like emojis for critical warnings. diff --git a/.gitattributes b/.gitattributes deleted file mode 100644 index a6b56c9..0000000 --- a/.gitattributes +++ /dev/null @@ -1,5 +0,0 @@ -images/*.png filter=lfs diff=lfs merge=lfs -text -images/*.jpg filter=lfs diff=lfs merge=lfs -text -images/*.jpeg filter=lfs diff=lfs merge=lfs -text -images/*.webp filter=lfs diff=lfs merge=lfs -text -images/*.gif filter=lfs diff=lfs merge=lfs -text \ No newline at end of file diff --git a/.gitignore b/.gitignore index 7a60b85..4239e62 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,4 @@ __pycache__/ *.pyc +Errors.md +Claude_Last_Convo.md \ No newline at end of file diff --git a/.jules/bolt.md b/.jules/bolt.md new file mode 100644 index 0000000..2d22201 --- /dev/null +++ b/.jules/bolt.md @@ -0,0 +1,3 @@ +## 2024-03-24 - Image Encoding Performance +**Learning:** `cv2.imencode` is ~4x faster than PIL's `save` for PNG images, even with low compression levels. However, PIL is ~30% faster for JPEG and avoids the overhead of converting PIL images back to Numpy/OpenCV format. +**Action:** Use a hybrid approach: Stick to OpenCV for PNG encoding, but use PIL directly for JPEG/WebP to save memory and CPU cycles on conversion. 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..4dfa41c --- /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..d5e6f0a --- /dev/null +++ b/bot/bot.py @@ -0,0 +1,127 @@ +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) + + + import uuid + self.client_id = str(uuid.uuid4()) + + # ComfyUI Clients + self.comfy_client = ComfyUIClient(base_url=config.comfyui.url) + self.comfy_ws = ComfyUIWebSocket(base_url=config.comfyui.url, client_id=self.client_id) + + # Services + from .services.delivery import DeliveryService + from .services.job_manager import JobManager + from .services.permissions import PermissionService + + self.delivery_service = DeliveryService(self, self.comfy_client) + self.job_manager = JobManager( + self.repository, + self.comfy_client, + self.comfy_ws, + self.delivery_service + ) + self.permission_service = PermissionService(self.repository) + + 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/admin.py b/bot/cogs/admin.py new file mode 100644 index 0000000..a2f9a34 --- /dev/null +++ b/bot/cogs/admin.py @@ -0,0 +1,68 @@ +import discord +from discord import app_commands +from discord.ext import commands +import logging + +from ..services.permissions import require_permission, Permissions, PermissionLevel + +logger = logging.getLogger(__name__) + +class AdminCog(commands.Cog): + def __init__(self, bot): + self.bot = bot + + @app_commands.command(name="admin", description="Bot configuration (Admin only)") + @require_permission(Permissions.ADMIN.value) + @app_commands.choices(action=[ + app_commands.Choice(name="status", value="status"), + ]) + async def admin(self, interaction: discord.Interaction, action: app_commands.Choice[str]): + if action.value == "status": + await self._show_status(interaction) + + async def _show_status(self, interaction: discord.Interaction): + comfy_status = await self.bot.comfy_client.check_status() + status_emoji = "✅" if comfy_status else "❌" + + embed = discord.Embed(title="Bot Status", color=discord.Color.dark_grey()) + embed.add_field(name="ComfyUI Connection", value=f"{status_emoji} {self.bot.config.comfyui_url}", inline=False) + embed.add_field(name="Guilds", value=str(len(self.bot.guilds)), inline=True) + embed.add_field(name="Latency", value=f"{round(self.bot.latency * 1000)}ms", inline=True) + + await interaction.response.send_message(embed=embed, ephemeral=True) + + @app_commands.command(name="setrole", description="Set permission level for a role") + @require_permission(Permissions.ADMIN.value) + @app_commands.choices(level=[ + app_commands.Choice(name="User", value="user"), + app_commands.Choice(name="Generator", value="generator"), + app_commands.Choice(name="Admin", value="admin"), + ]) + async def setrole(self, interaction: discord.Interaction, role: discord.Role, level: app_commands.Choice[str]): + """Assign a permission level to a Discord role.""" + if not interaction.guild: + await interaction.response.send_message("This command must be used in a server.", ephemeral=True) + return + + try: + # Ensure server exists in DB + await self.bot.repository.get_or_create_server( + str(interaction.guild.id), + interaction.guild.name + ) + + await self.bot.repository.set_server_role( + server_discord_id=str(interaction.guild.id), + role_discord_id=str(role.id), + permission_level=level.value + ) + await interaction.response.send_message( + f"✅ Role {role.mention} set to **{level.name}** permission level.", + ephemeral=True + ) + except Exception as e: + logger.error(f"Failed to set role: {e}") + await interaction.response.send_message("❌ Failed to update role permissions.", ephemeral=True) + +async def setup(bot): + await bot.add_cog(AdminCog(bot)) diff --git a/bot/cogs/generate.py b/bot/cogs/generate.py new file mode 100644 index 0000000..9ff4748 --- /dev/null +++ b/bot/cogs/generate.py @@ -0,0 +1,120 @@ +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 ..services.permissions import require_permission, Permissions +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") + ]) + @require_permission(Permissions.USER.value) + 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/cogs/queue.py b/bot/cogs/queue.py new file mode 100644 index 0000000..6941ab9 --- /dev/null +++ b/bot/cogs/queue.py @@ -0,0 +1,115 @@ +import discord +from discord import app_commands +from discord.ext import commands +import logging +from typing import Optional + +from ..services.permissions import require_permission, Permissions +from ..database.models import JobStatus + +logger = logging.getLogger(__name__) + +class QueueCog(commands.Cog): + def __init__(self, bot): + self.bot = bot + + @app_commands.command(name="queue", description="Manage the generation queue") + @app_commands.choices(action=[ + app_commands.Choice(name="view", value="view"), + app_commands.Choice(name="clear", value="clear"), + ]) + async def queue(self, interaction: discord.Interaction, action: app_commands.Choice[str]): + """General queue commands.""" + command = action.value + + if command == "view": + await self._view_queue(interaction) + elif command == "clear": + await self._clear_queue(interaction) + + async def _view_queue(self, interaction: discord.Interaction): + """Show current pending jobs.""" + await interaction.response.defer(ephemeral=True) + + pending_jobs = await self.bot.repository.get_pending_jobs() + + if not pending_jobs: + await interaction.followup.send("🟢 The queue is currently empty.", ephemeral=True) + return + + embed = discord.Embed( + title=f"Generation Queue ({len(pending_jobs)})", + color=discord.Color.blue() + ) + + # Show top 10 + desc_lines = [] + for i, job in enumerate(pending_jobs[:10]): + status_icon = "🔄" if job.status == JobStatus.RUNNING.value else "⏳" + user_mention = f"<@{job.user.discord_id}>" + prompt_text = job.positive_prompt or "No prompt provided" + prompt_preview = (prompt_text[:40] + "...") if len(prompt_text) > 40 else prompt_text + desc_lines.append(f"`#{i+1}` {status_icon} **ID:{job.id}** {user_mention}: {prompt_preview}") + + if len(pending_jobs) > 10: + desc_lines.append(f"...and {len(pending_jobs) - 10} more.") + + embed.description = "\n".join(desc_lines) + await interaction.followup.send(embed=embed, ephemeral=True) + + async def _clear_queue(self, interaction: discord.Interaction): + """Clear user's own pending jobs.""" + # Check permission manually since it's a subcommand handler + has_perm = await self.bot.permission_service.check_permission(interaction.user, Permissions.GENERATOR.value) + if not has_perm: + await interaction.response.send_message("⛔ You need the **Generator** role to clear the queue.", ephemeral=True) + return + + await interaction.response.defer(ephemeral=True) + + # Get user's pending jobs + # Note: repo.get_pending_jobs returns ALL. Better to filter or add new repo method. + # But for now let's iterate. + all_pending = await self.bot.repository.get_pending_jobs() + user_jobs = [j for j in all_pending if str(j.user.discord_id) == str(interaction.user.id)] + + count = 0 + for job in user_jobs: + success = await self.bot.job_manager.cancel_job(job.id) + if success: + count += 1 + + if count > 0: + await interaction.followup.send(f"🗑️ Cancelled {count} of your pending jobs.", ephemeral=True) + else: + await interaction.followup.send("No pending jobs found to clear.", ephemeral=True) + + + @app_commands.command(name="cancel", description="Cancel a specific job") + @app_commands.describe(job_id="The ID of the job to cancel") + @require_permission(Permissions.GENERATOR.value) + async def cancel(self, interaction: discord.Interaction, job_id: int): + """Cancel a specific job by ID.""" + await interaction.response.defer(ephemeral=True) + + job = await self.bot.repository.get_job_by_id(job_id) + if not job: + await interaction.followup.send(f"❌ Job ID {job_id} not found.", ephemeral=True) + return + + # Check permissions + is_owner = str(job.user.discord_id) == str(interaction.user.id) + is_admin = await self.bot.permission_service.check_permission(interaction.user, Permissions.ADMIN.value) + + if not is_owner and not is_admin: + await interaction.followup.send("⛔ You can only cancel your own jobs.", ephemeral=True) + return + + success = await self.bot.job_manager.cancel_job(job_id) + if success: + await interaction.followup.send(f"✅ Job {job_id} cancelled.", ephemeral=True) + else: + await interaction.followup.send(f"⚠️ Could not cancel job {job_id} (maybe already finished?).", ephemeral=True) + +async def setup(bot): + await bot.add_cog(QueueCog(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..612eb12 --- /dev/null +++ b/bot/comfyui/websocket.py @@ -0,0 +1,97 @@ +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" + self.client_id = client_id + 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}") + if self.session and not self.session.closed: + await self.session.close() + 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..cadfe9b --- /dev/null +++ b/bot/database/repository.py @@ -0,0 +1,656 @@ +""" +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 and username != "Unknown": + 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) + .options(selectinload(Job.user)) + .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: + current_images = [] + if job.output_images: + try: + current_images = json.loads(job.output_images) + except json.JSONDecodeError: + current_images = [] + + # Append new images + current_images.extend(output_images) + job.output_images = json.dumps(current_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..36b166d --- /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.duration:.1f}s" if hasattr(job, 'duration') and job.duration 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..352cbdb --- /dev/null +++ b/bot/services/job_manager.py @@ -0,0 +1,206 @@ +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 + # Use common client_id from WebSocket + self.client_id = comfy_ws.client_id + + # In-memory mapping of prompt_id -> current status buffer + self._active_jobs = {} + self._delivery_tasks = {} # prompt_id -> asyncio.Task + + 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.""" + + # 0. Validate and ensure entities exist + user = await self.repo.get_or_create_user(user_discord_id, "Unknown") + + # Check queue limit + current_queue_count = await self.repo.count_user_pending_jobs(user_discord_id, server_discord_id) + + # Get server specific limit or default + max_queue = 3 + if server_discord_id: + server = await self.repo.get_server(server_discord_id) + if server: + max_queue = server.max_queue_per_user + + if current_queue_count >= max_queue: + raise ValueError(f"Queue limit reached ({max_queue} jobs). Please wait for your current jobs to finish.") + + # 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 + + async def _schedule_delivery(self, job: Job): + """Schedule a debounced delivery for a job.""" + prompt_id = job.prompt_id + + # Cancel existing timer if any + if prompt_id in self._delivery_tasks: + self._delivery_tasks[prompt_id].cancel() + + # Create new timer + task = asyncio.create_task(self._deliver_delayed(job)) + self._delivery_tasks[prompt_id] = task + + # Cleanup callback + def cleanup(t): + if prompt_id in self._delivery_tasks and self._delivery_tasks[prompt_id] == t: + self._delivery_tasks.pop(prompt_id, None) + + task.add_done_callback(cleanup) + + async def _deliver_delayed(self, job: Job): + """Wait briefly then deliver the job.""" + try: + await asyncio.sleep(1.0) # Debounce window + logger.info(f"Job {job.id} delivery timer expired. Delivering results...") + await self.delivery.deliver_job(job) + except asyncio.CancelledError: + logger.debug(f"Job {job.id} delivery debounced/cancelled.") + except Exception as e: + logger.error(f"Error in delayed delivery for job {job.id}: {e}") + + # -- 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} update received. Scheduling delivery...") + await self._schedule_delivery(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/bot/services/permissions.py b/bot/services/permissions.py new file mode 100644 index 0000000..da69b84 --- /dev/null +++ b/bot/services/permissions.py @@ -0,0 +1,95 @@ +import logging +from typing import Optional, List, Dict, Union +import discord +from enum import Enum + +from ..database.repository import Repository +from ..database.models import PermissionLevel, ServerRole, User + +logger = logging.getLogger(__name__) + +class Permissions(Enum): + USER = "user" + GENERATOR = "generator" + ADMIN = "admin" + +class PermissionService: + """Manages role-based permissions.""" + + def __init__(self, repository: Repository): + self.repo = repository + + def get_permission_hierarchy(self) -> Dict[str, int]: + return { + Permissions.USER.value: 1, + Permissions.GENERATOR.value: 2, + Permissions.ADMIN.value: 3 + } + + + + async def get_user_permission_level(self, member: Union[discord.Member, discord.User]) -> str: + """ + Determine the highest permission level for a user in a guild. + Default is USER. + Administrator permission in Discord implies ADMIN level. + """ + # Handle DMs or non-guild context where we have User instead of Member + if isinstance(member, discord.User): + return Permissions.USER.value + + if member.guild_permissions.administrator: + return Permissions.ADMIN.value + + # Fetch configured roles for this server + server_roles = await self.repo.get_server_roles(str(member.guild.id)) + + if not server_roles: + return Permissions.USER.value + + hierarchy = self.get_permission_hierarchy() + current_level = Permissions.USER.value + current_score = hierarchy[current_level] + + # Check user's roles against configured roles + member_role_ids = [str(r.id) for r in member.roles] + + for server_role in server_roles: + if server_role.role_discord_id in member_role_ids: + level = server_role.permission_level + if level in hierarchy and hierarchy[level] > current_score: + current_level = level + current_score = hierarchy[level] + + return current_level + + async def check_permission(self, member: discord.Member, required_level: str) -> bool: + """Check if user meets the required permission level.""" + user_level = await self.get_user_permission_level(member) + hierarchy = self.get_permission_hierarchy() + + return hierarchy.get(user_level, 0) >= hierarchy.get(required_level, 0) + +# Helper decorator for checking permissions in commands +def require_permission(level: str): + async def predicate(interaction: discord.Interaction): + if not interaction.guild: + await interaction.response.send_message("⛔ This command cannot be used in DMs.", ephemeral=True) + return False + + + # We need to access the bot instance to get the permission service + bot = interaction.client + if not hasattr(bot, "permission_service"): + logger.error("Bot instance missing permission_service") + return False + + has_perm = await bot.permission_service.check_permission(interaction.user, level) + + if not has_perm: + await interaction.response.send_message( + f"⛔ You need **{level.upper()}** permission to use this command.", + ephemeral=True + ) + return has_perm + return discord.app_commands.check(predicate) diff --git a/discord_image_node.py b/discord_image_node.py index 2a00106..569be57 100644 --- a/discord_image_node.py +++ b/discord_image_node.py @@ -17,7 +17,7 @@ from uuid import uuid4 from typing import Any, Union, List, Optional # Import shared utilities -from utils import ( +from discordsend_utils import ( sanitize_json_for_export, update_github_cdn_urls, extract_prompts_from_workflow, @@ -63,7 +63,7 @@ class DiscordSendSaveImage: "min": 1, "max": 100, "step": 1, - "tooltip": "Quality for JPEG/WebP formats (1-100). Higher is better quality but larger file size." + "tooltip": "Quality (1-100) for JPEG/WebP. Ignored for PNG. Higher values = better quality but larger file size." }), "lossless": ("BOOLEAN", { "default": True, @@ -104,12 +104,12 @@ class DiscordSendSaveImage: "webhook_url": ("STRING", { "default": "", "multiline": False, - "tooltip": "Discord webhook URL to send images to. Leave empty to disable Discord integration." + "tooltip": "Discord webhook URL (from Server Settings > Integrations > Webhooks). Leave empty to disable Discord integration." }), "discord_message": ("STRING", { "default": "", "multiline": True, - "tooltip": "Optional message to send with the Discord images." + "tooltip": "Optional text to display with the image. Supports Discord Markdown (bold, italic, etc.)." }), "include_prompts_in_message": ("BOOLEAN", { "default": False, @@ -580,54 +580,53 @@ class DiscordSendSaveImage: # Send to Discord if enabled if send_to_discord and webhook_url: try: - # Prepare the image for Discord - use the resized PIL image (img) instead of original tensor - img_cv = np.array(img) - - # Convert RGB (PIL) to BGR (OpenCV) if needed - if len(img_cv.shape) == 3 and img_cv.shape[2] == 3: - img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGB2BGR) - - # Handle color conversion for special cases - if len(img_cv.shape) == 2: # Grayscale - img_cv = cv2.cvtColor(img_cv, cv2.COLOR_GRAY2BGR) - elif len(img_cv.shape) == 3 and img_cv.shape[2] == 4: # RGBA - img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGBA2BGRA) - # Generate unique filename for Discord using the selected format discord_filename = f"{uuid4()}.{file_format}" + file_bytes = BytesIO() + + # Optimization: Use PIL directly for JPEG/WebP to avoid numpy conversion overhead + # Use CV2 for PNG as it is significantly faster for that format - # Encode image using the selected format if file_format == "png": + # Use CV2 for PNG + img_cv = np.array(img) + + # Convert RGB (PIL) to BGR (OpenCV) + if len(img_cv.shape) == 3 and img_cv.shape[2] == 3: + img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGB2BGR) + + # Handle color conversion for special cases + if len(img_cv.shape) == 2: # Grayscale + img_cv = cv2.cvtColor(img_cv, cv2.COLOR_GRAY2BGR) + elif len(img_cv.shape) == 3 and img_cv.shape[2] == 4: # RGBA + img_cv = cv2.cvtColor(img_cv, cv2.COLOR_RGBA2BGRA) + _, buffer = cv2.imencode('.png', img_cv) + file_bytes = BytesIO(buffer) + elif file_format == "jpeg": - # JPEG is always lossy, but we can set quality to maximum if lossless is requested jpeg_quality = 100 if lossless else quality - encode_params = [int(cv2.IMWRITE_JPEG_QUALITY), jpeg_quality] - _, buffer = cv2.imencode('.jpg', img_cv, encode_params) + # JPEG does not support RGBA, convert to RGB if needed + save_img = img + if save_img.mode == 'RGBA': + save_img = save_img.convert('RGB') + save_img.save(file_bytes, format="JPEG", quality=jpeg_quality) + file_bytes.seek(0) + elif file_format == "webp": try: if lossless: - # For lossless WebP - using explicit parameter value as the constant may not be defined - # cv2.IMWRITE_WEBP_LOSSLESS is 9 in OpenCV - encode_params = [int(cv2.IMWRITE_WEBP_QUALITY), 100] # First ensure high quality - encode_params.extend([9, 1]) # 9 is the parameter ID for WEBP_LOSSLESS, 1 means true - _, buffer = cv2.imencode('.webp', img_cv, encode_params) - - # If that fails, try alternative method - if buffer is None or len(buffer) == 0: - raise ValueError("WebP lossless encoding failed with direct method") + img.save(file_bytes, format="WEBP", lossless=True) else: - # For lossy WebP with quality parameter - encode_params = [int(cv2.IMWRITE_WEBP_QUALITY), quality] - _, buffer = cv2.imencode('.webp', img_cv, encode_params) + img.save(file_bytes, format="WEBP", quality=quality) + file_bytes.seek(0) except Exception as e: print(f"Error with WebP encoding for Discord: {e}, falling back to PNG") - # Fallback to PNG if WebP encoding fails - _, buffer = cv2.imencode('.png', img_cv) - # Update filename to reflect the format change + # Fallback to PNG if WebP encoding fails (using PIL) discord_filename = f"{os.path.splitext(discord_filename)[0]}.png" - - file_bytes = BytesIO(buffer) + file_bytes = BytesIO() # Reset buffer + img.save(file_bytes, format="PNG", compress_level=self.compress_level) + file_bytes.seek(0) # If batch grouping is enabled, store the files for later if group_batched_images: diff --git a/discord_video_node.py b/discord_video_node.py index 909d5a4..8159894 100644 --- a/discord_video_node.py +++ b/discord_video_node.py @@ -23,7 +23,7 @@ import functools import server # Import shared utilities -from utils import sanitize_json_for_export, update_github_cdn_urls, send_to_discord_with_retry +from discordsend_utils import sanitize_json_for_export, update_github_cdn_urls, send_to_discord_with_retry # Define cached decorator for local use def cached(max_size=None): """ @@ -48,7 +48,7 @@ def cached(max_size=None): # Try to import dependencies from nodes.py try: - from utils import ffmpeg_path, get_audio, hash_path, validate_path, requeue_workflow, \ + from discordsend_utils import ffmpeg_path, get_audio, hash_path, validate_path, requeue_workflow, \ gifski_path, calculate_file_hash, strip_path, try_download_video, is_url, \ imageOrLatent, BIGMAX, merge_filter_args, ENCODE_ARGS, floatOrInt from comfy.utils import ProgressBar @@ -235,7 +235,7 @@ class DiscordSendSaveVideo: }), "add_time": ("BOOLEAN", { "default": True, - "tooltip": "Add the current time (HH-MM-SS) to filenames. Do not disable when sending videos to Discord." + "tooltip": "Add the current time (HH-MM-SS) to filenames. ⚠️ Recommended for Discord to avoid caching issues." }), "add_dimensions": ("BOOLEAN", { "default": True, @@ -250,12 +250,12 @@ class DiscordSendSaveVideo: "webhook_url": ("STRING", { "default": "", "multiline": False, - "tooltip": "Discord webhook URL to send videos to. Leave empty to disable Discord integration." + "tooltip": "Discord webhook URL (from Server Settings > Integrations > Webhooks). Leave empty to disable Discord integration." }), "discord_message": ("STRING", { "default": "", "multiline": True, - "tooltip": "Optional message to send with the Discord videos." + "tooltip": "Optional text to display with the video. Supports Discord Markdown (bold, italic, etc.)." }), "include_prompts_in_message": ("BOOLEAN", { "default": False, diff --git a/utils/__init__.py b/discordsend_utils/__init__.py similarity index 100% rename from utils/__init__.py rename to discordsend_utils/__init__.py diff --git a/utils/discord_api.py b/discordsend_utils/discord_api.py similarity index 99% rename from utils/discord_api.py rename to discordsend_utils/discord_api.py index 98e5abc..0523609 100644 --- a/utils/discord_api.py +++ b/discordsend_utils/discord_api.py @@ -11,6 +11,7 @@ import logging from io import BytesIO from typing import Any, Dict, List, Optional, Tuple +import json import requests # Get logger for this module @@ -216,7 +217,7 @@ class DiscordWebhookClient: if files: response = requests.post( self.webhook_url, - data={"payload_json": str(data)} if data else None, + data={"payload_json": json.dumps(data)} if data else None, files=files, timeout=60 ) @@ -408,4 +409,3 @@ def send_to_discord_with_retry( # Return the last response even if it was an error return response - diff --git a/utils/github_integration.py b/discordsend_utils/github_integration.py similarity index 100% rename from utils/github_integration.py rename to discordsend_utils/github_integration.py diff --git a/utils/logging_config.py b/discordsend_utils/logging_config.py similarity index 100% rename from utils/logging_config.py rename to discordsend_utils/logging_config.py diff --git a/utils/prompt_extractor.py b/discordsend_utils/prompt_extractor.py similarity index 100% rename from utils/prompt_extractor.py rename to discordsend_utils/prompt_extractor.py diff --git a/utils/sanitizer.py b/discordsend_utils/sanitizer.py similarity index 100% rename from utils/sanitizer.py rename to discordsend_utils/sanitizer.py diff --git a/discordsend_utils/workflow_builder.py b/discordsend_utils/workflow_builder.py new file mode 100644 index 0000000..09f26d9 --- /dev/null +++ b/discordsend_utils/workflow_builder.py @@ -0,0 +1,146 @@ +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 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/scripts/auto_review.sh b/scripts/auto_review.sh new file mode 100755 index 0000000..d1694bf --- /dev/null +++ b/scripts/auto_review.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +# auto_review.sh - Automatically review the latest changes using Gemini CLI + +# 1. Get the latest changes (staged or last commit) +# If there are staged changes, review them. Otherwise review last commit. +if git diff --quiet --cached; then + # No staged changes, check last commit + DIFF_CONTENT=$(git show HEAD) + CONTEXT="Review the following changes from the last commit:" +else + # Staged changes exist + DIFF_CONTENT=$(git diff --cached) + CONTEXT="Review the following staged changes:" +fi + +if [ -z "$DIFF_CONTENT" ]; then + echo "No changes found to review." + exit 0 +fi + +# 2. Construct Prompt +PROMPT="You are a Senior Software Engineer acting as a code reviewer. +Review the following code changes for: +1. Potential bugs or race conditions +2. Security vulnerabilities +3. Code style and best practices +4. Logical errors + +Be concise and constructive. + +$CONTEXT +\`\`\`diff +$DIFF_CONTENT +\`\`\` +" + +# 3. Call Gemini CLI +echo "🤖 Asking Gemini to review changes..." +echo "----------------------------------------" +gemini "$PROMPT" +echo "----------------------------------------" 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..08f94e2 --- /dev/null +++ b/tests/bot/test_imports.py @@ -0,0 +1,48 @@ +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.services.permissions + print("✅ bot.services.permissions imported") + + import bot.cogs.queue + print("✅ bot.cogs.queue imported") + + import bot.cogs.admin + print("✅ bot.cogs.admin 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/tests/test_utils.py b/tests/test_utils.py index 4546dc3..e80252d 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -12,8 +12,8 @@ import unittest # Add parent directory to path for imports sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from utils.sanitizer import sanitize_json_for_export -from utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging, send_to_discord_with_retry +from discordsend_utils.sanitizer import sanitize_json_for_export +from discordsend_utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging, send_to_discord_with_retry from unittest.mock import patch, MagicMock @@ -145,7 +145,7 @@ class TestSSRFPrevention(unittest.TestCase): self.assertIn("Invalid webhook URL", str(cm.exception)) - @patch('requests.post') + @patch('discordsend_utils.discord_api.requests.post') def test_send_to_discord_allows_valid_url(self, mock_post): """Should allow valid Discord URLs.""" valid_url = "https://discord.com/api/webhooks/123/abc"