feat: Implement Phase 1 & 2 of ComfyUI Companion Bot
This commit is contained in:
+361
@@ -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*
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -0,0 +1,42 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import sys
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# Add project root to python path to allow imports
|
||||
project_root = Path(__file__).resolve().parent.parent
|
||||
sys.path.append(str(project_root))
|
||||
|
||||
from bot.config import Config
|
||||
from bot.bot import ComfyUIBot
|
||||
from utils.logging_config import setup_logging
|
||||
|
||||
def main():
|
||||
# Setup logging
|
||||
setup_logging()
|
||||
logger = logging.getLogger("bot")
|
||||
|
||||
# Load configuration
|
||||
try:
|
||||
config = Config()
|
||||
except Exception as e:
|
||||
logger.critical(f"Failed to load configuration: {e}")
|
||||
return
|
||||
|
||||
if not config.discord_token:
|
||||
logger.critical("Discord token not found! Set DISCORDBOT_DISCORD_TOKEN env var or config.yaml")
|
||||
return
|
||||
|
||||
# Initialize and run bot
|
||||
bot = ComfyUIBot(config)
|
||||
|
||||
try:
|
||||
bot.run(config.discord_token)
|
||||
except KeyboardInterrupt:
|
||||
logger.info("Bot stopped by user.")
|
||||
except Exception as e:
|
||||
logger.critical(f"Bot crashed: {e}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+122
@@ -0,0 +1,122 @@
|
||||
import discord
|
||||
from discord.ext import commands
|
||||
import logging
|
||||
import sys
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
from .config import Config
|
||||
from .database.repository import Repository
|
||||
from .comfyui.client import ComfyUIClient
|
||||
from .comfyui.websocket import ComfyUIWebSocket
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class ComfyUIBot(commands.Bot):
|
||||
"""
|
||||
Main Bot Class for ComfyUI Companion.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
intents = discord.Intents.default()
|
||||
intents.message_content = True # Needed for some commands if not pure slash
|
||||
intents.members = True # Useful for permission checks
|
||||
|
||||
super().__init__(
|
||||
command_prefix=commands.when_mentioned,
|
||||
intents=intents,
|
||||
help_command=None,
|
||||
description="ComfyUI Companion Bot"
|
||||
)
|
||||
self.config = config
|
||||
|
||||
# Database
|
||||
self.repository = Repository(config.database_url)
|
||||
|
||||
|
||||
# ComfyUI Clients
|
||||
self.comfy_client = ComfyUIClient(base_url=config.comfyui_url)
|
||||
self.comfy_ws = ComfyUIWebSocket(base_url=config.comfyui_url)
|
||||
|
||||
# Services
|
||||
from .services.delivery import DeliveryService
|
||||
from .services.job_manager import JobManager
|
||||
|
||||
self.delivery_service = DeliveryService(self, self.comfy_client)
|
||||
self.job_manager = JobManager(
|
||||
self.repository,
|
||||
self.comfy_client,
|
||||
self.comfy_ws,
|
||||
self.delivery_service
|
||||
)
|
||||
|
||||
async def setup_hook(self):
|
||||
"""Async setup before bot starts."""
|
||||
logger.info(f"Setting up bot for {self.user}...")
|
||||
|
||||
# Initialize Database
|
||||
await self.repository.init_db()
|
||||
logger.info("Database initialized.")
|
||||
|
||||
# Connect to ComfyUI WebSocket
|
||||
try:
|
||||
await self.comfy_ws.connect()
|
||||
await self.job_manager.start()
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not connect to ComfyUI WebSocket at start: {e}")
|
||||
# We don't crash, as ComfyUI might come up later
|
||||
|
||||
# Load Cogs
|
||||
await self._load_cogs()
|
||||
|
||||
# Sync Commands
|
||||
# In production, sync might be manual or per-guild to avoid rate limits
|
||||
# For simplicity in this self-hosted bot, we'll sync global
|
||||
try:
|
||||
synced = await self.tree.sync()
|
||||
logger.info(f"Synced {len(synced)} slash commands.")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to sync commands: {e}")
|
||||
|
||||
async def _load_cogs(self):
|
||||
"""Load extensions from cogs directory."""
|
||||
# We need to use dotted path: bot.cogs.generate
|
||||
cogs_dir = Path(__file__).parent / "cogs"
|
||||
|
||||
# Extensions to load
|
||||
extensions = [
|
||||
"bot.cogs.generate",
|
||||
# "bot.cogs.queue",
|
||||
# "bot.cogs.templates",
|
||||
# "bot.cogs.history",
|
||||
# "bot.cogs.admin",
|
||||
]
|
||||
|
||||
for ext in extensions:
|
||||
try:
|
||||
# Check if file exists first to avoid confusing errors if we haven't created it yet
|
||||
# (Since we are building incrementally)
|
||||
module_name = ext.split(".")[-1]
|
||||
if (cogs_dir / f"{module_name}.py").exists():
|
||||
await self.load_extension(ext)
|
||||
logger.info(f"Loaded extension: {ext}")
|
||||
else:
|
||||
logger.debug(f"Skipping extension {ext} (file not found)")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load extension {ext}: {e}")
|
||||
|
||||
async def close(self):
|
||||
"""Cleanup on shutdown."""
|
||||
logger.info("Shutting down bot...")
|
||||
await self.comfy_client.close()
|
||||
await self.comfy_ws.disconnect()
|
||||
await self.repository.close()
|
||||
await super().close()
|
||||
|
||||
async def on_ready(self):
|
||||
logger.info(f"Bot logged in as {self.user} (ID: {self.user.id})")
|
||||
logger.info(f"Connected to {len(self.guilds)} guilds")
|
||||
|
||||
async def on_command_error(self, ctx, error):
|
||||
"""Global error handler for prefix commands (if any)."""
|
||||
logger.error(f"Command error: {error}", exc_info=False)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -0,0 +1,118 @@
|
||||
import discord
|
||||
from discord import app_commands
|
||||
from discord.ext import commands
|
||||
import logging
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from ..embeds.builders import EmbedBuilder
|
||||
from ...utils.workflow_builder import WorkflowBuilder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class GenerateCog(commands.Cog):
|
||||
def __init__(self, bot):
|
||||
self.bot = bot
|
||||
|
||||
@app_commands.command(name="generate", description="Generate an image using ComfyUI")
|
||||
@app_commands.describe(
|
||||
prompt="The positive prompt to generate",
|
||||
negative_prompt="Aspects to avoid (optional)",
|
||||
seed="Seed for generation (optional)",
|
||||
steps="Number of steps (optional)",
|
||||
cfg="CFG Scale (optional)",
|
||||
delivery="Delivery method (channel or dm)"
|
||||
)
|
||||
@app_commands.choices(delivery=[
|
||||
app_commands.Choice(name="Current Channel", value="channel"),
|
||||
app_commands.Choice(name="Direct Message", value="dm")
|
||||
])
|
||||
async def generate(self, interaction: discord.Interaction,
|
||||
prompt: str,
|
||||
negative_prompt: str = "",
|
||||
seed: int = None,
|
||||
steps: int = None,
|
||||
cfg: float = None,
|
||||
delivery: app_commands.Choice[str] = None):
|
||||
|
||||
await interaction.response.defer()
|
||||
|
||||
# Load default workflow
|
||||
# TODO: Move this path to config or database
|
||||
workflow_path = Path(__file__).parent.parent / "data" / "default_workflow_api.json"
|
||||
|
||||
if not workflow_path.exists():
|
||||
await interaction.followup.send("❌ Error: Default workflow not found.", ephemeral=True)
|
||||
return
|
||||
|
||||
try:
|
||||
with open(workflow_path, "r") as f:
|
||||
workflow_json = json.load(f)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load workflow: {e}")
|
||||
await interaction.followup.send("❌ Error: Failed to load workflow configuration.", ephemeral=True)
|
||||
return
|
||||
|
||||
# Prepare parameters for tracking
|
||||
parameters = {
|
||||
"seed": seed,
|
||||
"steps": steps,
|
||||
"cfg": cfg
|
||||
}
|
||||
|
||||
# Modify workflow
|
||||
builder = WorkflowBuilder(workflow_json)
|
||||
builder.set_prompt(prompt, negative_prompt)
|
||||
|
||||
if seed is not None:
|
||||
builder.set_seed(seed)
|
||||
else:
|
||||
# Random seed if not provided
|
||||
import random
|
||||
generated_seed = random.randint(1, 1000000000000000)
|
||||
builder.set_seed(generated_seed)
|
||||
parameters["seed"] = generated_seed # Track actual seed
|
||||
|
||||
if steps is not None:
|
||||
builder.set_steps(steps)
|
||||
|
||||
if cfg is not None:
|
||||
builder.set_cfg(cfg)
|
||||
|
||||
final_workflow = builder.get_workflow()
|
||||
|
||||
# Determine delivery method
|
||||
delivery_method = delivery.value if delivery else "channel"
|
||||
|
||||
# Check server context
|
||||
server_id = str(interaction.guild_id) if interaction.guild else None
|
||||
channel_id = str(interaction.channel_id)
|
||||
|
||||
try:
|
||||
# Create Job
|
||||
job = await self.bot.job_manager.create_job(
|
||||
user_discord_id=str(interaction.user.id),
|
||||
workflow=final_workflow,
|
||||
positive_prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
parameters=parameters,
|
||||
server_discord_id=server_id,
|
||||
channel_id=channel_id,
|
||||
delivery_type=delivery_method
|
||||
)
|
||||
|
||||
# Send Queued Embed
|
||||
embed = EmbedBuilder.job_queued(job)
|
||||
await interaction.followup.send(embed=embed)
|
||||
|
||||
# Store the interaction message ID if we want to update it later
|
||||
# (JobManager could use this to update the specific message)
|
||||
original_message = await interaction.original_response()
|
||||
await self.bot.repository.update_job_message(job.prompt_id, str(original_message.id))
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to start generation: {e}")
|
||||
await interaction.followup.send(f"❌ Error starting generation: {str(e)}", ephemeral=True)
|
||||
|
||||
async def setup(bot):
|
||||
await bot.add_cog(GenerateCog(bot))
|
||||
@@ -0,0 +1,6 @@
|
||||
"""ComfyUI integration package."""
|
||||
|
||||
from .client import ComfyUIClient
|
||||
from .websocket import ComfyUIWebSocket
|
||||
|
||||
__all__ = ["ComfyUIClient", "ComfyUIWebSocket"]
|
||||
@@ -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()
|
||||
@@ -0,0 +1,94 @@
|
||||
import aiohttp
|
||||
import logging
|
||||
import json
|
||||
import asyncio
|
||||
from typing import Callable, Coroutine, Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class ComfyUIWebSocket:
|
||||
"""WebSocket Client for real-time ComfyUI events."""
|
||||
|
||||
def __init__(self, base_url: str = "ws://127.0.0.1:8188", client_id: str = ""):
|
||||
# Convert http/https to ws/wss if needed
|
||||
if base_url.startswith("http://"):
|
||||
base_url = base_url.replace("http://", "ws://")
|
||||
elif base_url.startswith("https://"):
|
||||
base_url = base_url.replace("https://", "wss://")
|
||||
|
||||
self.ws_url = f"{base_url.rstrip('/')}/ws"
|
||||
if client_id:
|
||||
self.ws_url += f"?clientId={client_id}"
|
||||
|
||||
self.ws: Optional[aiohttp.ClientWebSocketResponse] = None
|
||||
self.session: Optional[aiohttp.ClientSession] = None
|
||||
self._callbacks: Dict[str, List[Callable[[Dict[str, Any]], Coroutine[Any, Any, None]]]] = {}
|
||||
self._running = False
|
||||
self._listen_task: Optional[asyncio.Task] = None
|
||||
|
||||
async def connect(self):
|
||||
"""Connect to the WebSocket."""
|
||||
if self.session is None or self.session.closed:
|
||||
self.session = aiohttp.ClientSession()
|
||||
|
||||
try:
|
||||
self.ws = await self.session.ws_connect(self.ws_url)
|
||||
self._running = True
|
||||
self._listen_task = asyncio.create_task(self._listen())
|
||||
logger.info(f"Connected to ComfyUI WebSocket at {self.ws_url}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to WebSocket: {e}")
|
||||
raise
|
||||
|
||||
async def disconnect(self):
|
||||
"""Disconnect from WebSocket."""
|
||||
self._running = False
|
||||
if self.ws:
|
||||
await self.ws.close()
|
||||
if self.session:
|
||||
await self.session.close()
|
||||
if self._listen_task:
|
||||
try:
|
||||
await self._listen_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
def add_listener(self, event_type: str, callback: Callable[[Dict[str, Any]], Coroutine[Any, Any, None]]):
|
||||
"""Register a callback for an event type."""
|
||||
if event_type not in self._callbacks:
|
||||
self._callbacks[event_type] = []
|
||||
self._callbacks[event_type].append(callback)
|
||||
|
||||
async def _listen(self):
|
||||
"""Listen loop for incoming messages."""
|
||||
if not self.ws:
|
||||
return
|
||||
|
||||
try:
|
||||
async for msg in self.ws:
|
||||
if msg.type == aiohttp.WSMsgType.TEXT:
|
||||
try:
|
||||
data = json.loads(msg.data)
|
||||
event_type = data.get("type", "unknown")
|
||||
# Some messages pack content in 'data', others at top level
|
||||
# ComfyUI typically sends {type: "event_name", data: {...}, sid: "..."}
|
||||
|
||||
handlers = self._callbacks.get(event_type, [])
|
||||
for handler in handlers:
|
||||
try:
|
||||
await handler(data)
|
||||
except Exception as e:
|
||||
logger.error(f"Error in WebSocket handler for {event_type}: {e}")
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Received invalid JSON: {msg.data}")
|
||||
elif msg.type == aiohttp.WSMsgType.ERROR:
|
||||
logger.error("WebSocket connection closed with error")
|
||||
break
|
||||
except Exception as e:
|
||||
if self._running:
|
||||
logger.error(f"WebSocket listener error: {e}")
|
||||
# Verify reconnection logic would handle this or let job manager handle it
|
||||
finally:
|
||||
if self._running:
|
||||
logger.info("WebSocket listener stopped unexpectedly.")
|
||||
+212
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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",
|
||||
]
|
||||
@@ -0,0 +1,222 @@
|
||||
"""
|
||||
SQLAlchemy ORM models for bot persistence.
|
||||
|
||||
Tables:
|
||||
- users: Discord user preferences
|
||||
- servers: Guild configuration
|
||||
- server_roles: Permission role mappings
|
||||
- templates: Saved prompt presets
|
||||
- jobs: Generation history and queue
|
||||
- workflows: Default workflow storage
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
Column,
|
||||
DateTime,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase, relationship
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
"""Base class for all models."""
|
||||
pass
|
||||
|
||||
|
||||
class DeliveryType(str, Enum):
|
||||
"""Delivery method for generated images."""
|
||||
CHANNEL = "channel"
|
||||
DM = "dm"
|
||||
|
||||
|
||||
class JobStatus(str, Enum):
|
||||
"""Status of a generation job."""
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class PermissionLevel(str, Enum):
|
||||
"""Permission levels for role-based access."""
|
||||
USER = "user"
|
||||
GENERATOR = "generator"
|
||||
ADMIN = "admin"
|
||||
|
||||
|
||||
class User(Base):
|
||||
"""Discord user preferences and data."""
|
||||
__tablename__ = "users"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
discord_id = Column(String(20), unique=True, nullable=False, index=True)
|
||||
username = Column(String(100), nullable=False)
|
||||
default_delivery = Column(String(10), default=DeliveryType.CHANNEL.value)
|
||||
default_workflow_id = Column(Integer, ForeignKey("workflows.id"), nullable=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
# Relationships
|
||||
templates = relationship("Template", back_populates="user", cascade="all, delete-orphan")
|
||||
jobs = relationship("Job", back_populates="user", cascade="all, delete-orphan")
|
||||
default_workflow = relationship("Workflow", foreign_keys=[default_workflow_id])
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<User(id={self.id}, discord_id={self.discord_id}, username={self.username})>"
|
||||
|
||||
|
||||
class Server(Base):
|
||||
"""Discord server (guild) configuration."""
|
||||
__tablename__ = "servers"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
discord_id = Column(String(20), unique=True, nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
default_channel_id = Column(String(20), nullable=True)
|
||||
enabled = Column(Boolean, default=True)
|
||||
max_queue_per_user = Column(Integer, default=3)
|
||||
default_workflow_id = Column(Integer, ForeignKey("workflows.id"), nullable=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
# Relationships
|
||||
roles = relationship("ServerRole", back_populates="server", cascade="all, delete-orphan")
|
||||
templates = relationship("Template", back_populates="server", cascade="all, delete-orphan")
|
||||
jobs = relationship("Job", back_populates="server")
|
||||
default_workflow = relationship("Workflow", foreign_keys=[default_workflow_id])
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Server(id={self.id}, discord_id={self.discord_id}, name={self.name})>"
|
||||
|
||||
|
||||
class ServerRole(Base):
|
||||
"""Role-based permission mapping for servers."""
|
||||
__tablename__ = "server_roles"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
server_id = Column(Integer, ForeignKey("servers.id"), nullable=False)
|
||||
role_discord_id = Column(String(20), nullable=False)
|
||||
permission_level = Column(String(20), nullable=False)
|
||||
|
||||
# Relationships
|
||||
server = relationship("Server", back_populates="roles")
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("server_id", "role_discord_id", name="uq_server_role"),
|
||||
Index("idx_server_roles_server", "server_id"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<ServerRole(server_id={self.server_id}, role={self.role_discord_id}, level={self.permission_level})>"
|
||||
|
||||
|
||||
class Workflow(Base):
|
||||
"""Stored workflow configurations."""
|
||||
__tablename__ = "workflows"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
workflow_json = Column(Text, nullable=False)
|
||||
is_default = Column(Boolean, default=False)
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Workflow(id={self.id}, name={self.name}, is_default={self.is_default})>"
|
||||
|
||||
|
||||
class Template(Base):
|
||||
"""User-saved prompt templates."""
|
||||
__tablename__ = "templates"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
|
||||
server_id = Column(Integer, ForeignKey("servers.id"), nullable=True) # NULL = private
|
||||
name = Column(String(100), nullable=False)
|
||||
positive_prompt = Column(Text, nullable=False)
|
||||
negative_prompt = Column(Text, default="")
|
||||
parameters = Column(Text, nullable=True) # JSON: {steps, cfg, seed, etc.}
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
# Relationships
|
||||
user = relationship("User", back_populates="templates")
|
||||
server = relationship("Server", back_populates="templates")
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_templates_user", "user_id"),
|
||||
Index("idx_templates_server", "server_id"),
|
||||
UniqueConstraint("user_id", "server_id", "name", name="uq_template_name"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Template(id={self.id}, name={self.name}, user_id={self.user_id})>"
|
||||
|
||||
|
||||
class Job(Base):
|
||||
"""Generation job tracking and history."""
|
||||
__tablename__ = "jobs"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
prompt_id = Column(String(64), unique=True, nullable=False, index=True)
|
||||
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
|
||||
server_id = Column(Integer, ForeignKey("servers.id"), nullable=True)
|
||||
channel_id = Column(String(20), nullable=True)
|
||||
|
||||
# Status tracking
|
||||
status = Column(String(20), default=JobStatus.PENDING.value, index=True)
|
||||
queue_position = Column(Integer, nullable=True)
|
||||
progress = Column(Integer, default=0)
|
||||
progress_max = Column(Integer, default=0)
|
||||
|
||||
# Generation parameters
|
||||
positive_prompt = Column(Text, nullable=True)
|
||||
negative_prompt = Column(Text, nullable=True)
|
||||
parameters = Column(Text, nullable=True) # JSON
|
||||
workflow_json = Column(Text, nullable=True)
|
||||
|
||||
# Results
|
||||
output_images = Column(Text, nullable=True) # JSON array
|
||||
error_message = Column(Text, nullable=True)
|
||||
|
||||
# Delivery
|
||||
delivery_type = Column(String(10), default=DeliveryType.CHANNEL.value)
|
||||
message_id = Column(String(20), nullable=True) # Discord message for updates
|
||||
|
||||
# Timestamps
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
started_at = Column(DateTime, nullable=True)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
|
||||
# Relationships
|
||||
user = relationship("User", back_populates="jobs")
|
||||
server = relationship("Server", back_populates="jobs")
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_jobs_user", "user_id"),
|
||||
Index("idx_jobs_status", "status"),
|
||||
Index("idx_jobs_created", "created_at"),
|
||||
)
|
||||
|
||||
@property
|
||||
def duration(self) -> Optional[float]:
|
||||
"""Get job duration in seconds."""
|
||||
if self.started_at and self.completed_at:
|
||||
return (self.completed_at - self.started_at).total_seconds()
|
||||
return None
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Job(id={self.id}, prompt_id={self.prompt_id}, status={self.status})>"
|
||||
@@ -0,0 +1,645 @@
|
||||
"""
|
||||
Data access layer for bot database operations.
|
||||
|
||||
Provides async CRUD operations for all database models.
|
||||
"""
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select, update, delete, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from .models import (
|
||||
Base,
|
||||
User,
|
||||
Server,
|
||||
ServerRole,
|
||||
Template,
|
||||
Job,
|
||||
Workflow,
|
||||
JobStatus,
|
||||
PermissionLevel,
|
||||
)
|
||||
|
||||
|
||||
class Repository:
|
||||
"""Async repository for database operations."""
|
||||
|
||||
def __init__(self, database_url: str):
|
||||
"""
|
||||
Initialize the repository.
|
||||
|
||||
Args:
|
||||
database_url: SQLAlchemy async database URL
|
||||
"""
|
||||
self.engine = create_async_engine(database_url, echo=False)
|
||||
self.async_session = async_sessionmaker(
|
||||
self.engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
)
|
||||
|
||||
async def init_db(self) -> None:
|
||||
"""Create all database tables."""
|
||||
async with self.engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the database connection."""
|
||||
await self.engine.dispose()
|
||||
|
||||
# ==================== User Operations ====================
|
||||
|
||||
async def get_or_create_user(
|
||||
self,
|
||||
discord_id: str,
|
||||
username: str,
|
||||
) -> User:
|
||||
"""Get existing user or create new one."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(User).where(User.discord_id == discord_id)
|
||||
)
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
if user is None:
|
||||
user = User(discord_id=discord_id, username=username)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
elif user.username != username:
|
||||
user.username = username
|
||||
await session.commit()
|
||||
|
||||
return user
|
||||
|
||||
async def get_user(self, discord_id: str) -> Optional[User]:
|
||||
"""Get user by Discord ID."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(User).where(User.discord_id == discord_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def update_user_delivery(
|
||||
self,
|
||||
discord_id: str,
|
||||
delivery_type: str,
|
||||
) -> Optional[User]:
|
||||
"""Update user's default delivery preference."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(User).where(User.discord_id == discord_id)
|
||||
)
|
||||
user = result.scalar_one_or_none()
|
||||
if user:
|
||||
user.default_delivery = delivery_type
|
||||
await session.commit()
|
||||
return user
|
||||
|
||||
# ==================== Server Operations ====================
|
||||
|
||||
async def get_or_create_server(
|
||||
self,
|
||||
discord_id: str,
|
||||
name: str,
|
||||
) -> Server:
|
||||
"""Get existing server or create new one."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Server).where(Server.discord_id == discord_id)
|
||||
)
|
||||
server = result.scalar_one_or_none()
|
||||
|
||||
if server is None:
|
||||
server = Server(discord_id=discord_id, name=name)
|
||||
session.add(server)
|
||||
await session.commit()
|
||||
await session.refresh(server)
|
||||
elif server.name != name:
|
||||
server.name = name
|
||||
await session.commit()
|
||||
|
||||
return server
|
||||
|
||||
async def get_server(self, discord_id: str) -> Optional[Server]:
|
||||
"""Get server by Discord ID."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Server)
|
||||
.options(selectinload(Server.roles))
|
||||
.where(Server.discord_id == discord_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def update_server_channel(
|
||||
self,
|
||||
discord_id: str,
|
||||
channel_id: str,
|
||||
) -> Optional[Server]:
|
||||
"""Update server's default output channel."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Server).where(Server.discord_id == discord_id)
|
||||
)
|
||||
server = result.scalar_one_or_none()
|
||||
if server:
|
||||
server.default_channel_id = channel_id
|
||||
await session.commit()
|
||||
return server
|
||||
|
||||
async def update_server_queue_limit(
|
||||
self,
|
||||
discord_id: str,
|
||||
limit: int,
|
||||
) -> Optional[Server]:
|
||||
"""Update server's per-user queue limit."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Server).where(Server.discord_id == discord_id)
|
||||
)
|
||||
server = result.scalar_one_or_none()
|
||||
if server:
|
||||
server.max_queue_per_user = limit
|
||||
await session.commit()
|
||||
return server
|
||||
|
||||
# ==================== Role Operations ====================
|
||||
|
||||
async def set_server_role(
|
||||
self,
|
||||
server_discord_id: str,
|
||||
role_discord_id: str,
|
||||
permission_level: str,
|
||||
) -> ServerRole:
|
||||
"""Set or update a role's permission level for a server."""
|
||||
async with self.async_session() as session:
|
||||
# Get server
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise ValueError(f"Server {server_discord_id} not found")
|
||||
|
||||
# Check for existing role mapping
|
||||
role_result = await session.execute(
|
||||
select(ServerRole).where(
|
||||
ServerRole.server_id == server.id,
|
||||
ServerRole.role_discord_id == role_discord_id,
|
||||
)
|
||||
)
|
||||
role = role_result.scalar_one_or_none()
|
||||
|
||||
if role:
|
||||
role.permission_level = permission_level
|
||||
else:
|
||||
role = ServerRole(
|
||||
server_id=server.id,
|
||||
role_discord_id=role_discord_id,
|
||||
permission_level=permission_level,
|
||||
)
|
||||
session.add(role)
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(role)
|
||||
return role
|
||||
|
||||
async def get_server_roles(self, server_discord_id: str) -> list[ServerRole]:
|
||||
"""Get all role mappings for a server."""
|
||||
async with self.async_session() as session:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if not server:
|
||||
return []
|
||||
|
||||
result = await session.execute(
|
||||
select(ServerRole).where(ServerRole.server_id == server.id)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def delete_server_role(
|
||||
self,
|
||||
server_discord_id: str,
|
||||
role_discord_id: str,
|
||||
) -> bool:
|
||||
"""Remove a role mapping."""
|
||||
async with self.async_session() as session:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if not server:
|
||||
return False
|
||||
|
||||
result = await session.execute(
|
||||
delete(ServerRole).where(
|
||||
ServerRole.server_id == server.id,
|
||||
ServerRole.role_discord_id == role_discord_id,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
return result.rowcount > 0
|
||||
|
||||
# ==================== Template Operations ====================
|
||||
|
||||
async def create_template(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
name: str,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str = "",
|
||||
parameters: Optional[dict] = None,
|
||||
server_discord_id: Optional[str] = None,
|
||||
) -> Template:
|
||||
"""Create a new prompt template."""
|
||||
async with self.async_session() as session:
|
||||
# Get user
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
raise ValueError(f"User {user_discord_id} not found")
|
||||
|
||||
# Get server if provided
|
||||
server_id = None
|
||||
if server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
server_id = server.id
|
||||
|
||||
template = Template(
|
||||
user_id=user.id,
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
positive_prompt=positive_prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
parameters=json.dumps(parameters) if parameters else None,
|
||||
)
|
||||
session.add(template)
|
||||
await session.commit()
|
||||
await session.refresh(template)
|
||||
return template
|
||||
|
||||
async def get_template(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
name: str,
|
||||
server_discord_id: Optional[str] = None,
|
||||
) -> Optional[Template]:
|
||||
"""Get a template by name."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
return None
|
||||
|
||||
# Build query
|
||||
query = select(Template).where(
|
||||
Template.user_id == user.id,
|
||||
Template.name == name,
|
||||
)
|
||||
|
||||
if server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
query = query.where(Template.server_id == server.id)
|
||||
else:
|
||||
query = query.where(Template.server_id.is_(None))
|
||||
|
||||
result = await session.execute(query)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_templates(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
server_discord_id: Optional[str] = None,
|
||||
include_shared: bool = True,
|
||||
) -> list[Template]:
|
||||
"""List templates for a user."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
return []
|
||||
|
||||
# Private templates
|
||||
query = select(Template).where(
|
||||
Template.user_id == user.id,
|
||||
Template.server_id.is_(None),
|
||||
)
|
||||
result = await session.execute(query)
|
||||
templates = list(result.scalars().all())
|
||||
|
||||
# Shared templates in server
|
||||
if include_shared and server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
shared_query = select(Template).where(
|
||||
Template.server_id == server.id
|
||||
)
|
||||
shared_result = await session.execute(shared_query)
|
||||
templates.extend(shared_result.scalars().all())
|
||||
|
||||
return templates
|
||||
|
||||
async def delete_template(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
name: str,
|
||||
server_discord_id: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Delete a template."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
return False
|
||||
|
||||
query = delete(Template).where(
|
||||
Template.user_id == user.id,
|
||||
Template.name == name,
|
||||
)
|
||||
|
||||
if server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
query = query.where(Template.server_id == server.id)
|
||||
else:
|
||||
query = query.where(Template.server_id.is_(None))
|
||||
|
||||
result = await session.execute(query)
|
||||
await session.commit()
|
||||
return result.rowcount > 0
|
||||
|
||||
# ==================== Job Operations ====================
|
||||
|
||||
async def create_job(
|
||||
self,
|
||||
prompt_id: str,
|
||||
user_discord_id: str,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str = "",
|
||||
parameters: Optional[dict] = None,
|
||||
workflow_json: Optional[str] = None,
|
||||
delivery_type: str = "channel",
|
||||
server_discord_id: Optional[str] = None,
|
||||
channel_id: Optional[str] = None,
|
||||
) -> Job:
|
||||
"""Create a new generation job."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
raise ValueError(f"User {user_discord_id} not found")
|
||||
|
||||
server_id = None
|
||||
if server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
server_id = server.id
|
||||
|
||||
job = Job(
|
||||
prompt_id=prompt_id,
|
||||
user_id=user.id,
|
||||
server_id=server_id,
|
||||
channel_id=channel_id,
|
||||
positive_prompt=positive_prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
parameters=json.dumps(parameters) if parameters else None,
|
||||
workflow_json=workflow_json,
|
||||
delivery_type=delivery_type,
|
||||
)
|
||||
session.add(job)
|
||||
await session.commit()
|
||||
await session.refresh(job)
|
||||
return job
|
||||
|
||||
async def get_job(self, prompt_id: str) -> Optional[Job]:
|
||||
"""Get a job by prompt ID."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job)
|
||||
.options(selectinload(Job.user))
|
||||
.where(Job.prompt_id == prompt_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_job_by_id(self, job_id: int) -> Optional[Job]:
|
||||
"""Get a job by internal ID."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job)
|
||||
.options(selectinload(Job.user))
|
||||
.where(Job.id == job_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def update_job_status(
|
||||
self,
|
||||
prompt_id: str,
|
||||
status: str,
|
||||
error_message: Optional[str] = None,
|
||||
output_images: Optional[list[str]] = None,
|
||||
) -> Optional[Job]:
|
||||
"""Update job status."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job).where(Job.prompt_id == prompt_id)
|
||||
)
|
||||
job = result.scalar_one_or_none()
|
||||
if not job:
|
||||
return None
|
||||
|
||||
job.status = status
|
||||
|
||||
if status == JobStatus.RUNNING.value:
|
||||
job.started_at = datetime.utcnow()
|
||||
elif status in (JobStatus.COMPLETED.value, JobStatus.FAILED.value, JobStatus.CANCELLED.value):
|
||||
job.completed_at = datetime.utcnow()
|
||||
|
||||
if error_message is not None:
|
||||
job.error_message = error_message
|
||||
|
||||
if output_images is not None:
|
||||
job.output_images = json.dumps(output_images)
|
||||
|
||||
await session.commit()
|
||||
return job
|
||||
|
||||
async def update_job_progress(
|
||||
self,
|
||||
prompt_id: str,
|
||||
progress: int,
|
||||
progress_max: int,
|
||||
) -> Optional[Job]:
|
||||
"""Update job progress."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job).where(Job.prompt_id == prompt_id)
|
||||
)
|
||||
job = result.scalar_one_or_none()
|
||||
if job:
|
||||
job.progress = progress
|
||||
job.progress_max = progress_max
|
||||
await session.commit()
|
||||
return job
|
||||
|
||||
async def update_job_message(
|
||||
self,
|
||||
prompt_id: str,
|
||||
message_id: str,
|
||||
) -> Optional[Job]:
|
||||
"""Update the Discord message ID for a job."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job).where(Job.prompt_id == prompt_id)
|
||||
)
|
||||
job = result.scalar_one_or_none()
|
||||
if job:
|
||||
job.message_id = message_id
|
||||
await session.commit()
|
||||
return job
|
||||
|
||||
async def list_user_jobs(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
limit: int = 10,
|
||||
status: Optional[str] = None,
|
||||
) -> list[Job]:
|
||||
"""List jobs for a user."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
return []
|
||||
|
||||
query = (
|
||||
select(Job)
|
||||
.where(Job.user_id == user.id)
|
||||
.order_by(Job.created_at.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
if status:
|
||||
query = query.where(Job.status == status)
|
||||
|
||||
result = await session.execute(query)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def count_user_pending_jobs(
|
||||
self,
|
||||
user_discord_id: str,
|
||||
server_discord_id: Optional[str] = None,
|
||||
) -> int:
|
||||
"""Count pending/running jobs for a user."""
|
||||
async with self.async_session() as session:
|
||||
user_result = await session.execute(
|
||||
select(User).where(User.discord_id == user_discord_id)
|
||||
)
|
||||
user = user_result.scalar_one_or_none()
|
||||
if not user:
|
||||
return 0
|
||||
|
||||
query = (
|
||||
select(func.count(Job.id))
|
||||
.where(Job.user_id == user.id)
|
||||
.where(Job.status.in_([JobStatus.PENDING.value, JobStatus.RUNNING.value]))
|
||||
)
|
||||
|
||||
if server_discord_id:
|
||||
server_result = await session.execute(
|
||||
select(Server).where(Server.discord_id == server_discord_id)
|
||||
)
|
||||
server = server_result.scalar_one_or_none()
|
||||
if server:
|
||||
query = query.where(Job.server_id == server.id)
|
||||
|
||||
result = await session.execute(query)
|
||||
return result.scalar() or 0
|
||||
|
||||
async def get_pending_jobs(self) -> list[Job]:
|
||||
"""Get all pending jobs ordered by creation time."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Job)
|
||||
.options(selectinload(Job.user))
|
||||
.where(Job.status.in_([JobStatus.PENDING.value, JobStatus.RUNNING.value]))
|
||||
.order_by(Job.created_at)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
# ==================== Workflow Operations ====================
|
||||
|
||||
async def save_workflow(
|
||||
self,
|
||||
name: str,
|
||||
workflow_json: str,
|
||||
description: Optional[str] = None,
|
||||
is_default: bool = False,
|
||||
) -> Workflow:
|
||||
"""Save a workflow configuration."""
|
||||
async with self.async_session() as session:
|
||||
# If setting as default, unset other defaults
|
||||
if is_default:
|
||||
await session.execute(
|
||||
update(Workflow).where(Workflow.is_default == True).values(is_default=False)
|
||||
)
|
||||
|
||||
workflow = Workflow(
|
||||
name=name,
|
||||
workflow_json=workflow_json,
|
||||
description=description,
|
||||
is_default=is_default,
|
||||
)
|
||||
session.add(workflow)
|
||||
await session.commit()
|
||||
await session.refresh(workflow)
|
||||
return workflow
|
||||
|
||||
async def get_default_workflow(self) -> Optional[Workflow]:
|
||||
"""Get the default workflow."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Workflow).where(Workflow.is_default == True)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_workflow(self, name: str) -> Optional[Workflow]:
|
||||
"""Get a workflow by name."""
|
||||
async with self.async_session() as session:
|
||||
result = await session.execute(
|
||||
select(Workflow).where(Workflow.name == name)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Discord embed builders."""
|
||||
|
||||
from .builders import EmbedBuilder
|
||||
|
||||
__all__ = ["EmbedBuilder"]
|
||||
@@ -0,0 +1,72 @@
|
||||
import discord
|
||||
import logging
|
||||
from typing import Optional, List
|
||||
from datetime import datetime
|
||||
|
||||
class EmbedBuilder:
|
||||
"""Helper for building Discord embeds."""
|
||||
|
||||
@staticmethod
|
||||
def job_queued(job, position: int = 0) -> discord.Embed:
|
||||
"""Embed for queued job."""
|
||||
embed = discord.Embed(
|
||||
title="🎨 Generation Queued",
|
||||
description=f"**Prompt:** {job.positive_prompt}",
|
||||
color=discord.Color.blue(),
|
||||
timestamp=datetime.utcnow()
|
||||
)
|
||||
embed.add_field(name="Queue Position", value=str(position) if position > 0 else "Pending...", inline=True)
|
||||
embed.add_field(name="Status", value="Waiting to start...", inline=True)
|
||||
if job.negative_prompt:
|
||||
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
|
||||
embed.set_footer(text=f"Job ID: {job.id}")
|
||||
return embed
|
||||
|
||||
@staticmethod
|
||||
def job_progress(job, progress: int, max_progress: int) -> discord.Embed:
|
||||
"""Embed for running job with progress."""
|
||||
percent = int((progress / max_progress) * 100) if max_progress > 0 else 0
|
||||
bars = "█" * (percent // 10) + "░" * (10 - (percent // 10))
|
||||
|
||||
embed = discord.Embed(
|
||||
title="🎨 Generating...",
|
||||
description=f"**Prompt:** {job.positive_prompt}",
|
||||
color=discord.Color.orange(),
|
||||
timestamp=datetime.utcnow()
|
||||
)
|
||||
embed.add_field(name="Progress", value=f"`{bars}` {percent}%", inline=False)
|
||||
if job.negative_prompt:
|
||||
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
|
||||
embed.set_footer(text=f"Job ID: {job.id}")
|
||||
return embed
|
||||
|
||||
@staticmethod
|
||||
def job_completed(job, image_count: int) -> discord.Embed:
|
||||
"""Embed for completed job."""
|
||||
embed = discord.Embed(
|
||||
title="✨ Generation Complete!",
|
||||
description=f"**Prompt:** {job.positive_prompt}",
|
||||
color=discord.Color.green(),
|
||||
timestamp=datetime.utcnow()
|
||||
)
|
||||
embed.add_field(name="Images", value=f"{image_count} generated", inline=True)
|
||||
embed.add_field(name="Duration", value=f"{job.execution_time:.1f}s" if hasattr(job, 'execution_time') and job.execution_time else "Done", inline=True)
|
||||
|
||||
if job.negative_prompt:
|
||||
embed.add_field(name="Negative Prompt", value=job.negative_prompt, inline=False)
|
||||
|
||||
embed.set_footer(text=f"Job ID: {job.id}")
|
||||
return embed
|
||||
|
||||
@staticmethod
|
||||
def job_failed(job, error_message: str) -> discord.Embed:
|
||||
"""Embed for failed job."""
|
||||
embed = discord.Embed(
|
||||
title="❌ Generation Failed",
|
||||
description=f"**Prompt:** {job.positive_prompt}",
|
||||
color=discord.Color.red(),
|
||||
timestamp=datetime.utcnow()
|
||||
)
|
||||
embed.add_field(name="Error", value=f"```{error_message}```", inline=False)
|
||||
embed.set_footer(text=f"Job ID: {job.id}")
|
||||
return embed
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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}")
|
||||
@@ -0,0 +1,158 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import json
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional, Dict, List, Any
|
||||
|
||||
from ..database.repository import Repository
|
||||
from ..database.models import JobStatus, Job
|
||||
from ..comfyui.client import ComfyUIClient
|
||||
from ..comfyui.websocket import ComfyUIWebSocket
|
||||
from .delivery import DeliveryService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class JobManager:
|
||||
"""Manages the lifecycle of generation jobs."""
|
||||
|
||||
def __init__(self, repository: Repository, comfy_client: ComfyUIClient, comfy_ws: ComfyUIWebSocket, delivery_service: DeliveryService):
|
||||
self.repo = repository
|
||||
self.client = comfy_client
|
||||
self.ws = comfy_ws
|
||||
self.delivery = delivery_service
|
||||
self.client_id = str(uuid.uuid4())
|
||||
|
||||
# In-memory mapping of prompt_id -> current status buffer
|
||||
self._active_jobs = {}
|
||||
|
||||
async def start(self):
|
||||
"""Start listening to WebSocket events."""
|
||||
self.ws.add_listener("status", self._on_status)
|
||||
self.ws.add_listener("execution_start", self._on_execution_start)
|
||||
self.ws.add_listener("executing", self._on_executing)
|
||||
self.ws.add_listener("executed", self._on_executed)
|
||||
self.ws.add_listener("execution_error", self._on_execution_error)
|
||||
self.ws.add_listener("progress", self._on_progress)
|
||||
logger.info(f"JobManager started with client_id: {self.client_id}")
|
||||
|
||||
async def create_job(self,
|
||||
user_discord_id: str,
|
||||
workflow: Dict[str, Any],
|
||||
positive_prompt: str,
|
||||
negative_prompt: str = "",
|
||||
parameters: Optional[Dict] = None,
|
||||
server_discord_id: Optional[str] = None,
|
||||
channel_id: Optional[str] = None,
|
||||
delivery_type: str = "channel") -> Job:
|
||||
"""Submit a job to ComfyUI and database."""
|
||||
|
||||
# 1. Submit to ComfyUI
|
||||
response = await self.client.queue_prompt(workflow, self.client_id)
|
||||
prompt_id = response.get("prompt_id")
|
||||
|
||||
if not prompt_id:
|
||||
raise ValueError("Failed to get prompt_id from ComfyUI")
|
||||
|
||||
# 2. Create DB entry
|
||||
job = await self.repo.create_job(
|
||||
prompt_id=prompt_id,
|
||||
user_discord_id=user_discord_id,
|
||||
server_discord_id=server_discord_id,
|
||||
channel_id=channel_id,
|
||||
positive_prompt=positive_prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
parameters=parameters,
|
||||
workflow_json=json.dumps(workflow),
|
||||
delivery_type=delivery_type
|
||||
)
|
||||
|
||||
logger.info(f"Created job {job.id} (prompt_id: {prompt_id}) for user {user_discord_id}")
|
||||
return job
|
||||
|
||||
async def cancel_job(self, job_id: int) -> bool:
|
||||
"""Cancel a job."""
|
||||
job = await self.repo.get_job_by_id(job_id)
|
||||
if not job:
|
||||
return False
|
||||
|
||||
if job.status in [JobStatus.COMPLETED.value, JobStatus.FAILED.value, JobStatus.CANCELLED.value]:
|
||||
return False
|
||||
|
||||
if job.status == JobStatus.RUNNING.value:
|
||||
await self.client.interrupt()
|
||||
|
||||
# Remove from queue if pending
|
||||
try:
|
||||
await self.client.delete_queue_item(job.prompt_id)
|
||||
except:
|
||||
pass
|
||||
|
||||
await self.repo.update_job_status(job.prompt_id, JobStatus.CANCELLED.value)
|
||||
return True
|
||||
|
||||
# -- WebSocket Event Handlers --
|
||||
|
||||
async def _on_status(self, data: Dict[str, Any]):
|
||||
pass
|
||||
|
||||
async def _on_execution_start(self, data: Dict[str, Any]):
|
||||
prompt_id = data.get("data", {}).get("prompt_id")
|
||||
if prompt_id:
|
||||
await self.repo.update_job_status(prompt_id, JobStatus.RUNNING.value)
|
||||
|
||||
async def _on_executing(self, data: Dict[str, Any]):
|
||||
pass
|
||||
|
||||
async def _on_progress(self, data: Dict[str, Any]):
|
||||
msg = data.get("data", {})
|
||||
prompt_id = msg.get("prompt_id")
|
||||
value = msg.get("value")
|
||||
max_val = msg.get("max")
|
||||
|
||||
if prompt_id and value is not None and max_val is not None:
|
||||
await self.repo.update_job_progress(prompt_id, value, max_val)
|
||||
|
||||
async def _on_executed(self, data: Dict[str, Any]):
|
||||
"""
|
||||
Handle execution completion of a node.
|
||||
If it contains images, we assume it's a relevant output.
|
||||
We accumulate images and update job status.
|
||||
"""
|
||||
msg = data.get("data", {})
|
||||
prompt_id = msg.get("prompt_id")
|
||||
output = msg.get("output", {})
|
||||
|
||||
if prompt_id and "images" in output:
|
||||
# Found images
|
||||
images = output["images"]
|
||||
|
||||
# Update DB with images and mark as completed
|
||||
# NOTE: In complex workflows, there might be multiple outputs.
|
||||
# Ideally we check if this is the last one or something.
|
||||
# But normally 'executed' with images means we got something.
|
||||
# We'll mark as completed for now. If multiple exist, latest wins.
|
||||
|
||||
job = await self.repo.update_job_status(
|
||||
prompt_id,
|
||||
JobStatus.COMPLETED.value,
|
||||
output_images=images
|
||||
)
|
||||
|
||||
if job:
|
||||
logger.info(f"Job {job.id} completed. Delivering results...")
|
||||
await self.delivery.deliver_job(job)
|
||||
|
||||
async def _on_execution_error(self, data: Dict[str, Any]):
|
||||
msg = data.get("data", {})
|
||||
prompt_id = msg.get("prompt_id")
|
||||
exception_type = msg.get("exception_type", "Unknown Error")
|
||||
exception_message = msg.get("exception_message", "")
|
||||
|
||||
if prompt_id:
|
||||
error_msg = f"{exception_type}: {exception_message}"
|
||||
job = await self.repo.update_job_status(prompt_id, JobStatus.FAILED.value, error_message=error_msg)
|
||||
# Notify user of failure via delivery service?
|
||||
# Ideally yes, but DeliveryService currently only sends images.
|
||||
# We might want to expand DeliveryService to handle errors too, or reuse the channel_id to post the failure embed.
|
||||
pass
|
||||
+7
-1
@@ -1 +1,7 @@
|
||||
requests>=2.25.0
|
||||
requests>=2.25.0
|
||||
discord.py>=2.3.0
|
||||
aiohttp>=3.9.0
|
||||
sqlalchemy>=2.0.0
|
||||
aiosqlite>=0.19.0
|
||||
pydantic>=2.0.0
|
||||
pyyaml>=6.0.0
|
||||
@@ -0,0 +1,39 @@
|
||||
import sys
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# Add project root to python path
|
||||
project_root = Path(__file__).resolve().parent.parent.parent
|
||||
sys.path.append(str(project_root))
|
||||
|
||||
print(f"Testing imports from {project_root}")
|
||||
|
||||
try:
|
||||
import bot.config
|
||||
print("✅ bot.config imported")
|
||||
|
||||
import bot.database.models
|
||||
print("✅ bot.database.models imported")
|
||||
|
||||
import bot.database.repository
|
||||
print("✅ bot.database.repository imported")
|
||||
|
||||
import bot.comfyui.client
|
||||
print("✅ bot.comfyui.client imported")
|
||||
|
||||
import bot.services.job_manager
|
||||
print("✅ bot.services.job_manager imported")
|
||||
|
||||
import bot.cogs.generate
|
||||
print("✅ bot.cogs.generate imported")
|
||||
|
||||
import bot.bot
|
||||
print("✅ bot.bot imported")
|
||||
|
||||
print("All modules imported successfully.")
|
||||
except ImportError as e:
|
||||
print(f"❌ ImportError: {e}")
|
||||
sys.exit(1)
|
||||
except Exception as e:
|
||||
print(f"❌ Error: {e}")
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,147 @@
|
||||
import json
|
||||
import random
|
||||
import logging
|
||||
from typing import Dict, Any, Optional, Tuple, List
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class WorkflowBuilder:
|
||||
"""Helper to manipulate ComfyUI workflow JSONs."""
|
||||
|
||||
def __init__(self, workflow_json: Dict[str, Any]):
|
||||
self.workflow = workflow_json
|
||||
# Create lookups
|
||||
self._nodes = self.workflow
|
||||
if "nodes" in self.workflow and isinstance(self.workflow["nodes"], list):
|
||||
# Handle "graph" format vs "api" format if needed.
|
||||
# But usually for API we stick to the {node_id: node_data} format.
|
||||
# If input is graph format, it might need conversion or distinct handling.
|
||||
# Assuming API format for now as that's what's sent to /prompt.
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def from_json_string(cls, json_str: str) -> 'WorkflowBuilder':
|
||||
"""Load from JSON int."""
|
||||
return cls(json.loads(json_str))
|
||||
|
||||
def get_workflow(self) -> Dict[str, Any]:
|
||||
"""Get the current workflow dict."""
|
||||
return self.workflow
|
||||
|
||||
def set_prompt(self, positive: str, negative: Optional[str] = None) -> None:
|
||||
"""
|
||||
Attempt to set positive and negative prompts.
|
||||
Heuristics:
|
||||
- Look for CLIPTextEncode nodes.
|
||||
- Often one is connected to KSampler 'positive' and one to 'negative'.
|
||||
- Or look for custom titles like 'Positive Prompt', 'Negative Prompt'.
|
||||
"""
|
||||
# Simple heuristic: Find CLIPTextEncode nodes
|
||||
# If we have title/coloring, we can use that.
|
||||
# Otherwise, we might need graph traversal to see what connects to KSampler.
|
||||
|
||||
# For MVP, let's assume standard ComfyUI structure or look for specific titles first
|
||||
|
||||
positive_node_id = self._find_node_by_title("Positive Prompt")
|
||||
negative_node_id = self._find_node_by_title("Negative Prompt")
|
||||
|
||||
# Fallback: Find KSampler and trace back
|
||||
if not positive_node_id or not negative_node_id:
|
||||
ksampler_id, ksampler = self._find_node_by_class("KSampler")
|
||||
if ksampler:
|
||||
# KSampler inputs: model, positive, negative, latent_image
|
||||
if not positive_node_id:
|
||||
positive_node_id = self._trace_input(ksampler, "positive")
|
||||
if not negative_node_id:
|
||||
negative_node_id = self._trace_input(ksampler, "negative")
|
||||
|
||||
if positive_node_id:
|
||||
self._update_node_input(positive_node_id, "text", positive)
|
||||
else:
|
||||
logger.warning("Could not identify Positive Prompt node.")
|
||||
|
||||
if negative and negative_node_id:
|
||||
self._update_node_input(negative_node_id, "text", negative)
|
||||
elif negative:
|
||||
logger.warning("Could not identify Negative Prompt node.")
|
||||
|
||||
def set_seed(self, seed: int) -> int:
|
||||
"""Set seed on KSampler nodes or Seed nodes."""
|
||||
# Find KSampler or anything with a 'seed' widget
|
||||
updated = False
|
||||
for node_id, node in self.workflow.items():
|
||||
if "inputs" in node:
|
||||
if "seed" in node["inputs"]:
|
||||
# Ensure it's an int widget, not a link
|
||||
if isinstance(node["inputs"]["seed"], (int, float)) or (isinstance(node["inputs"]["seed"], str) and node["inputs"]["seed"].isdigit()):
|
||||
node["inputs"]["seed"] = seed
|
||||
updated = True
|
||||
if "noise_seed" in node["inputs"]:
|
||||
# Some nodes call it noise_seed
|
||||
if isinstance(node["inputs"]["noise_seed"], (int, float)):
|
||||
node["inputs"]["noise_seed"] = seed
|
||||
updated = True
|
||||
|
||||
if not updated:
|
||||
logger.warning("Could not find any seed inputs to update.")
|
||||
return seed
|
||||
|
||||
def set_image_dimensions(self, width: int, height: int) -> None:
|
||||
"""Set width and height on EmptyLatentImage nodes."""
|
||||
node_id, _ = self._find_node_by_class("EmptyLatentImage")
|
||||
if node_id:
|
||||
self._update_node_input(node_id, "width", width)
|
||||
self._update_node_input(node_id, "height", height)
|
||||
|
||||
def set_steps(self, steps: int) -> None:
|
||||
"""Set steps on KSampler."""
|
||||
ksampler_ids = self._find_nodes_by_class("KSampler")
|
||||
for nid in ksampler_ids:
|
||||
self._update_node_input(nid, "steps", steps)
|
||||
|
||||
def set_cfg(self, cfg: float) -> None:
|
||||
"""Set CFG scale on KSampler."""
|
||||
ksampler_ids = self._find_nodes_by_class("KSampler")
|
||||
for nid in ksampler_ids:
|
||||
self._update_node_input(nid, "cfg", cfg)
|
||||
|
||||
def _find_node_by_title(self, title: str) -> Optional[str]:
|
||||
"""Find node by its custom title (`_meta.title`)."""
|
||||
for node_id, node in self.workflow.items():
|
||||
if "_meta" in node and node["_meta"].get("title") == title:
|
||||
return node_id
|
||||
return None
|
||||
|
||||
def _find_node_by_class(self, class_type: str) -> Tuple[Optional[str], Optional[Dict]]:
|
||||
"""Find first node of a specific class type."""
|
||||
for node_id, node in self.workflow.items():
|
||||
if node.get("class_type") == class_type:
|
||||
return node_id, node
|
||||
return None, None
|
||||
|
||||
def _find_nodes_by_class(self, class_type: str) -> List[str]:
|
||||
"""Find all nodes of a specific class type."""
|
||||
ids = []
|
||||
for node_id, node in self.workflow.items():
|
||||
if node.get("class_type") == class_type:
|
||||
ids.append(node_id)
|
||||
return ids
|
||||
|
||||
def _trace_input(self, node: Dict, input_name: str) -> Optional[str]:
|
||||
"""
|
||||
Trace back an input link to find the source node.
|
||||
Input format in API JSON: "input_name": ["source_node_id", slot_index]
|
||||
"""
|
||||
if "inputs" not in node or input_name not in node["inputs"]:
|
||||
return None
|
||||
|
||||
link = node["inputs"][input_name]
|
||||
# Link structure: [node_id, slot_idx]
|
||||
if isinstance(link, list) and len(link) == 2:
|
||||
return str(link[0])
|
||||
return None
|
||||
|
||||
def _update_node_input(self, node_id: str, input_name: str, value: Any) -> None:
|
||||
if node_id in self.workflow and "inputs" in self.workflow[node_id]:
|
||||
self.workflow[node_id]["inputs"][input_name] = value
|
||||
|
||||
Reference in New Issue
Block a user