Compare commits
+54
-12
@@ -1,23 +1,65 @@
|
||||
# Core Code API Configuration
|
||||
# Copy to .env and fill in real values
|
||||
|
||||
# Application settings
|
||||
# =============================================================================
|
||||
# Application
|
||||
# =============================================================================
|
||||
APP_NAME="Core Code API"
|
||||
APP_VERSION="1.0.0"
|
||||
DEBUG=false
|
||||
LOG_LEVEL=INFO
|
||||
|
||||
# Server settings
|
||||
# =============================================================================
|
||||
# Server
|
||||
# =============================================================================
|
||||
HOST=0.0.0.0
|
||||
PORT=8083
|
||||
|
||||
# CORS settings (default allows all origins for internal use)
|
||||
# CORS (default allows all origins for internal use)
|
||||
# CORS_ORIGINS=["http://192.168.86.149:82"]
|
||||
|
||||
# Logging
|
||||
LOG_LEVEL=INFO
|
||||
# =============================================================================
|
||||
# Infrastructure Services
|
||||
# =============================================================================
|
||||
|
||||
# Web Scraper Module
|
||||
WEB_SCRAPER_REQUEST_TIMEOUT=30
|
||||
WEB_SCRAPER_MAX_REDIRECTS=5
|
||||
WEB_SCRAPER_USER_AGENT="Mozilla/5.0 (compatible; CoreCode/1.0)"
|
||||
WEB_SCRAPER_DEFAULT_MAX_LENGTH=10000
|
||||
WEB_SCRAPER_MAX_LINKS_TO_EXTRACT=50
|
||||
# Portainer API (required for container/stack management)
|
||||
PORTAINER_URL=http://localhost:8001
|
||||
PORTAINER_API_KEY=ptr_your-api-key-here
|
||||
|
||||
# Nginx Proxy Manager API
|
||||
NPM_URL=http://localhost:81
|
||||
NPM_EMAIL=admin@example.com
|
||||
NPM_PASSWORD=your-npm-password
|
||||
|
||||
# =============================================================================
|
||||
# Home Automation
|
||||
# =============================================================================
|
||||
|
||||
# Home Assistant API
|
||||
HOMEASSISTANT_URL=http://localhost:8123
|
||||
HOMEASSISTANT_TOKEN=your-long-lived-access-token
|
||||
|
||||
# =============================================================================
|
||||
# AI Services
|
||||
# =============================================================================
|
||||
|
||||
# Ollama API
|
||||
OLLAMA_BASE_URL=http://localhost:11434
|
||||
|
||||
# SearXNG (self-hosted search)
|
||||
SEARXNG_URL=http://localhost:8080
|
||||
|
||||
# =============================================================================
|
||||
# Database (PostgreSQL)
|
||||
# =============================================================================
|
||||
|
||||
POSTGRES_HOST=localhost:5432
|
||||
POSTGRES_USER=core_api
|
||||
POSTGRES_PASSWORD=your-password
|
||||
|
||||
# =============================================================================
|
||||
# Vector Database
|
||||
# =============================================================================
|
||||
|
||||
# Qdrant
|
||||
QDRANT_HOST=qdrant
|
||||
QDRANT_PORT=6333
|
||||
@@ -1,12 +1,25 @@
|
||||
name: Build and Push
|
||||
|
||||
on:
|
||||
release:
|
||||
types: [published]
|
||||
push:
|
||||
tags:
|
||||
- 'v*'
|
||||
|
||||
jobs:
|
||||
release:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Create Gitea Release
|
||||
run: |
|
||||
curl -sf -X POST \
|
||||
-H "Authorization: token ${{ secrets.GITHUB_TOKEN }}" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"tag_name": "${{ github.ref_name }}", "name": "Release ${{ github.ref_name }}", "body": "Automated release for ${{ github.ref_name }}"}' \
|
||||
"${{ github.server_url }}/api/v1/repos/${{ github.repository }}/releases"
|
||||
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
needs: release
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
@@ -25,3 +38,9 @@ jobs:
|
||||
tags: |
|
||||
git.schweitz.internal/jpmschweitzer/core-api:latest
|
||||
git.schweitz.internal/jpmschweitzer/core-api:${{ github.ref_name }}
|
||||
|
||||
- name: Trigger Watchtower update
|
||||
if: success()
|
||||
run: |
|
||||
curl -sf -H "Authorization: Bearer ${{ secrets.WATCHTOWER_TOKEN }}" \
|
||||
http://watchtower:8080/v1/update
|
||||
|
||||
@@ -105,7 +105,6 @@ data/
|
||||
|
||||
# Credentials and secrets
|
||||
src/credentials.py
|
||||
credentials.py
|
||||
*.pem
|
||||
*.key
|
||||
secrets/
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
|
||||
# AGENTS.md
|
||||
|
||||
> **Start every session by reading this file.**
|
||||
> This file outlines the operational protocols, coding standards, and architectural decisions for this FastAPI project.
|
||||
|
||||
## 1. Agent Operational Protocols
|
||||
|
||||
### 🧠 Work Patterns (Plan-Act-Reflect)
|
||||
* **Plan:** Before writing code, briefly outline your plan. Identify which files you will touch and what the side effects might be.
|
||||
* **Act:** Execute the changes in small, atomic steps.
|
||||
* **Reflect:** After coding, verify your work. Did you break existing tests? Did you add new tests?
|
||||
|
||||
### 🛡️ Git Discipline
|
||||
* **NEVER commit to `main` or `master` directly.** Always create a feature branch: `feature/your-feature-name` or `fix/issue-description`.
|
||||
* **Commit Messages:** Use the [Conventional Commits](https://www.conventionalcommits.org/) format.
|
||||
* `feat: add user login endpoint`
|
||||
* `fix: resolve database connection timeout`
|
||||
* `refactor: split monolith dependency file`
|
||||
* **Atomic Commits:** Keep commits small. One logical change = one commit.
|
||||
|
||||
### 📝 Changelog Maintenance
|
||||
* **Update `CHANGELOG.md`** with every user-facing change.
|
||||
* Format: `## [Unreleased] - YYYY-MM-DD` followed by `### Added`, `### Changed`, or `### Fixed`.
|
||||
|
||||
### 🚀 Release Flow
|
||||
When changes are ready for deployment:
|
||||
|
||||
1. **Ask user if deploy cycle is desired **
|
||||
|
||||
2. **Update version** in `pyproject.toml`:
|
||||
- Bug fixes: bump patch version (1.8.3 → 1.8.4)
|
||||
- New features: bump minor version (1.8.4 → 1.9.0)
|
||||
|
||||
3. **Update CHANGELOG.md**:
|
||||
- Move items from `[Unreleased]` to new version section
|
||||
- Add release date: `## [1.8.4] - 2025-12-16`
|
||||
|
||||
4. **Commit and tag**:
|
||||
```bash
|
||||
git add -A
|
||||
git commit -m "fix: description of changes"
|
||||
git tag v1.8.4
|
||||
git push origin main --tags
|
||||
```
|
||||
|
||||
5. **CI/CD triggers automatically**:
|
||||
- Gitea CI builds Docker image on new tag
|
||||
- Watchtower pulls and deploys to production
|
||||
- Verify deployment: `curl http://192.168.86.149:8000/health`
|
||||
|
||||
---
|
||||
|
||||
## 2. FastAPI Architecture & Best Practices
|
||||
*Reference: [FastAPI Best Practices](https://github.com/zhanymkanov/fastapi-best-practices)*
|
||||
|
||||
### 📂 Project Structure (Directory-based, NOT File-type based)
|
||||
Do **not** group files by type (e.g., one huge `routers` folder). Group by **domain/module** inside a `src/` directory.
|
||||
|
||||
**Correct Structure:**
|
||||
```text
|
||||
src/
|
||||
├── auth/
|
||||
│ ├── router.py # Endpoints
|
||||
│ ├── schemas.py # Pydantic models
|
||||
│ ├── service.py # Business logic (CRUD, etc.)
|
||||
│ ├── dependencies.py# Module-specific dependencies
|
||||
│ └── config.py # Module-specific settings
|
||||
├── posts/
|
||||
│ ├── router.py
|
||||
│ └── ...
|
||||
└── main.py # App entry point
|
||||
+249
@@ -5,6 +5,255 @@ All notable changes to this project will be documented in this file.
|
||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
|
||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||
|
||||
## [1.9.0] - 2026-01-03
|
||||
|
||||
### Added
|
||||
|
||||
- **NPM Forward Auth Support** - Web authentication via Nginx Proxy Manager forward auth
|
||||
- `GET /auth/me` - Get current user from NPM forward auth headers (X-authentik-uid, X-authentik-email, etc.)
|
||||
- Auto-creates user on first web login if not in database
|
||||
- Syncs roles from NPM forward auth groups header
|
||||
- `get_user_by_email` and `get_user_by_authentik_id` methods in AuthService
|
||||
- Comprehensive tests for `/auth/me` endpoint
|
||||
|
||||
## [1.8.0] - 2026-01-03
|
||||
|
||||
### Added
|
||||
|
||||
- **Group-Role Mapping & Permissions** (Phase 3)
|
||||
- Decoupled group-role architecture (groups from Authentik, roles admin-managed)
|
||||
- Permission format: `domain.category:action` with action hierarchy
|
||||
- `require_permission` and `require_any_permission` dependency factories
|
||||
- Global admin override (`admin.general:admin`)
|
||||
- **User Profile & API Keys** (Phase 4)
|
||||
- `GET /auth/users/me` - Full user profile with roles and preferences
|
||||
- `GET/PATCH /auth/users/me/preferences` - User preferences management
|
||||
- `GET/POST/DELETE /auth/users/me/api-keys` - API key lifecycle
|
||||
- API keys with `tak_` prefix, SHA-256 hashing, shown only once on creation
|
||||
|
||||
## [1.7.0] - 2026-01-03
|
||||
|
||||
### Added
|
||||
|
||||
- **System Stats API** - Host system resource monitoring for dashboard widgets
|
||||
- `GET /tools/system/stats` - Real-time host system statistics
|
||||
- CPU: usage percentage, core count, load averages
|
||||
- Memory: usage percentage, total/used/available bytes
|
||||
- Disks: all mounted filesystems with usage stats (auto-discovers mounts)
|
||||
- Network: total bytes sent/received
|
||||
- GPU/VRAM: NVIDIA GPU memory usage (via nvidia-smi if available)
|
||||
- `psutil` dependency for cross-platform system metrics
|
||||
|
||||
## [1.6.1] - 2026-01-03
|
||||
|
||||
### Added
|
||||
|
||||
- `link_type` field to quick links for iframe vs new tab behavior
|
||||
|
||||
## [1.6.0] - 2026-01-03
|
||||
|
||||
### Added
|
||||
|
||||
- **Dashboard API** - Quick links and widgets management for Organizr-style dashboard
|
||||
- `GET /dashboard/quick-links` - List quick links with category/visibility filtering
|
||||
- `GET /dashboard/quick-links/{id}` - Get single quick link
|
||||
- `POST /dashboard/quick-links` - Create quick link
|
||||
- `PUT /dashboard/quick-links/{id}` - Update quick link
|
||||
- `DELETE /dashboard/quick-links/{id}` - Delete quick link
|
||||
- `POST /dashboard/quick-links/reorder` - Reorder quick links by position
|
||||
- `GET /dashboard/widgets` - List dashboard widgets
|
||||
- `GET /dashboard/widgets/{id}` - Get single widget
|
||||
- `POST /dashboard/widgets` - Create widget
|
||||
- `PUT /dashboard/widgets/{id}` - Update widget
|
||||
- `DELETE /dashboard/widgets/{id}` - Delete widget
|
||||
- Database migrations for `quick_links` and `dashboard_widgets` tables
|
||||
- Static file controller for serving Organizr widgets (`/static/widgets`)
|
||||
- Default local user authentication when OIDC is disabled
|
||||
|
||||
### Changed
|
||||
|
||||
- **Domain-based architecture** - Refactored codebase to domain-driven structure
|
||||
- `src/domains/` - Domain modules (auth, dashboard, health, housekeeping, infrastructure, tools)
|
||||
- `src/shared/` - Shared utilities (base, config, database, logging, security, clients)
|
||||
- Test suite updated for new domain structure (285 tests passing)
|
||||
|
||||
## [1.5.0] - 2026-01-01
|
||||
|
||||
### Added
|
||||
|
||||
- **Groups Management** - Authentik group synchronization
|
||||
- `GET /auth/groups` - List all groups with search and pagination
|
||||
- `POST /auth/groups/sync-from-authentik` - Bulk sync groups from Authentik admin API
|
||||
- Database model and migration for groups table
|
||||
|
||||
## [1.4.6] - 2026-01-01
|
||||
|
||||
### Changed
|
||||
|
||||
- Code cleanup: move inline `re` import to top of auth/service.py
|
||||
|
||||
## [1.4.5] - 2026-01-01
|
||||
|
||||
### Fixed
|
||||
|
||||
- Separate httpx and SQLAlchemy async contexts in bulk sync (fixes greenlet error)
|
||||
|
||||
## [1.4.4] - 2026-01-01
|
||||
|
||||
### Fixed
|
||||
|
||||
- Use `uuid` field instead of `pk` for Authentik user sync (pk is integer, uuid is proper UUID)
|
||||
- Skip internal_service_account type users during bulk sync
|
||||
|
||||
## [1.4.3] - 2026-01-01
|
||||
|
||||
### Fixed
|
||||
|
||||
- Manually extract and send session cookies for Authentik flow auth (fixes cross-domain cookie handling)
|
||||
|
||||
## [1.4.2] - 2026-01-01
|
||||
|
||||
### Fixed
|
||||
|
||||
- Use Authentik domain URL (auth.schweitz.net) instead of IP to fix cookie domain matching
|
||||
|
||||
## [1.4.1] - 2026-01-01
|
||||
|
||||
### Fixed
|
||||
|
||||
- Authentik config now uses AUTHENTIK_USERNAME/PASSWORD to match production env vars
|
||||
|
||||
## [1.4.0] - 2026-01-01
|
||||
|
||||
### Added
|
||||
|
||||
- **Authentication & User Management** - Authentik integration for user synchronization
|
||||
- `GET /auth/me` - Get current authenticated user info
|
||||
- `GET /auth/users` - List all users with search and pagination
|
||||
- `POST /auth/users/sync-from-authentik` - Bulk sync users from Authentik admin API
|
||||
- PostgreSQL database integration with async SQLAlchemy
|
||||
- Alembic database migrations for schema management
|
||||
- Database models: User, Role, UserPreferences, ApiKey
|
||||
- Token validation via Authentik userinfo endpoint
|
||||
- Role synchronization from Authentik groups
|
||||
- Health check now includes database connectivity status
|
||||
|
||||
### Changed
|
||||
|
||||
- Authentik configuration now uses username/password for admin API access
|
||||
- Health endpoint includes database status in diagnostics
|
||||
|
||||
## [1.3.1] - 2025-12-31
|
||||
|
||||
### Changed
|
||||
|
||||
- Simplified configuration: all settings now read from environment variables/.env only
|
||||
- Removed Docker socket fallback for container operations (Portainer API is now required)
|
||||
- Updated default search provider to SearXNG
|
||||
|
||||
### Removed
|
||||
|
||||
- `src/credentials.py` - credentials now managed via environment variables
|
||||
- Docker socket fallback methods from Portainer client
|
||||
- Unused search API settings (Brave, Google)
|
||||
|
||||
## [1.3.0] - 2025-12-31
|
||||
|
||||
### Added
|
||||
|
||||
- **Stack Management Endpoints** for Tatlock Control Room integration
|
||||
- `GET /infrastructure/stacks/{stackId}/compose` - Get stack Docker Compose YAML
|
||||
- `PUT /infrastructure/stacks/{stackId}/compose` - Update stack Docker Compose YAML
|
||||
- `GET /infrastructure/stacks/{stackId}/env` - Get stack environment variables
|
||||
- `PUT /infrastructure/stacks/{stackId}/env` - Update stack environment variables
|
||||
- `POST /infrastructure/stacks/{stackId}/deploy` - Redeploy stack
|
||||
- `POST /infrastructure/stacks/{stackId}/rebuild` - Pull images and recreate containers
|
||||
- `DELETE /infrastructure/containers/{id}` - Delete container with optional force flag
|
||||
- Portainer client methods: `get_stack_file`, `redeploy_stack`, `update_stack_env`, `delete_container`, `restart_container`
|
||||
|
||||
### Removed
|
||||
|
||||
- REQUESTED_SERVICES.md - endpoint specifications now implemented
|
||||
|
||||
## [1.2.1] - 2025-12-17
|
||||
|
||||
### Fixed
|
||||
|
||||
- Device control endpoint now returns correct new state after action (added 300ms delay for HA state propagation)
|
||||
|
||||
### Added
|
||||
|
||||
- AGENTS.md with project coding guidelines and release flow documentation
|
||||
|
||||
### Changed
|
||||
|
||||
- Removed obsolete web scraper settings from .env.example
|
||||
|
||||
## [1.2.0] - 2025-12-17
|
||||
|
||||
### Added
|
||||
|
||||
- **Housekeeping API** - Home Assistant integration for smart home control
|
||||
- `GET /housekeeping/health` - HA connection status
|
||||
- `GET /housekeeping/devices` - List controllable devices with optional domain/area filtering
|
||||
- `GET /housekeeping/devices/{entity_id}` - Get device details
|
||||
- `POST /housekeeping/devices/{entity_id}/control` - Control devices (turn_on, turn_off, toggle, set_brightness)
|
||||
- `GET /housekeeping/scenes` - List available scenes
|
||||
- `POST /housekeeping/scenes/{scene_id}/activate` - Activate a scene
|
||||
- `GET /housekeeping/scripts` - List available scripts
|
||||
- `POST /housekeeping/scripts/{script_id}/run` - Run a script
|
||||
- `GET /housekeeping/automations` - List automations
|
||||
- `POST /housekeeping/automations/{automation_id}/toggle` - Enable/disable automation
|
||||
- `GET /housekeeping/history` - Query state history
|
||||
- `GET /housekeeping/areas` - List rooms/areas
|
||||
- Home Assistant REST API client (`src/clients/homeassistant_client.py`)
|
||||
- Home Assistant configuration in credentials and settings
|
||||
- Comprehensive test suite with 65% code coverage (285 tests)
|
||||
- Tests for NPM client, Ollama client, AI client, OIDC authentication
|
||||
- Tests for infrastructure, health, tools, and housekeeping endpoints
|
||||
|
||||
### Removed
|
||||
|
||||
- **Web Scraper** - Entire web scraping module removed
|
||||
- `src/web_scraper/` directory deleted
|
||||
- `/web-scraper/scrape` endpoint removed
|
||||
- Trafilatura and BeautifulSoup dependencies removed from scraping use
|
||||
|
||||
### Changed
|
||||
|
||||
- Updated README.md with current architecture and all endpoints
|
||||
- Tools controller now only contains DNS lookup functionality
|
||||
- Health controller endpoints list updated to reflect current features
|
||||
|
||||
## [1.1.2] - 2024-12-14
|
||||
|
||||
### Added
|
||||
|
||||
- Version field to /health endpoint response
|
||||
|
||||
## [1.1.1] - 2024-12-14
|
||||
|
||||
### Added
|
||||
|
||||
- Watchtower trigger step in CI workflow for automatic container updates
|
||||
|
||||
## [1.1.0] - 2024-12-14
|
||||
|
||||
### Removed
|
||||
|
||||
- Uptime Kuma integration (kuma_client.py deleted)
|
||||
- All /infrastructure/monitors endpoints
|
||||
- Kuma monitor pause/resume from service start/stop operations
|
||||
- Uptime percentage display from service control widget
|
||||
- Kuma-related configuration and credentials
|
||||
- Stale memory tests (memory functionality moved to core-ai service)
|
||||
|
||||
### Changed
|
||||
|
||||
- Service control widget now uses Portainer container status exclusively
|
||||
- Simplified widget-data endpoint response (removed monitors field)
|
||||
- Updated service start/stop to only manage containers via Portainer
|
||||
|
||||
## [1.0.0] - 2024-12-14
|
||||
|
||||
### Added
|
||||
|
||||
@@ -1,33 +1,96 @@
|
||||
# Core Code API
|
||||
|
||||
OpenAPI-compatible functions for Open WebUI, providing web scraping and data processing capabilities.
|
||||
Central API service providing infrastructure management, home automation, and utility endpoints for the homelab ecosystem.
|
||||
|
||||
## Features
|
||||
|
||||
### Web Scraper
|
||||
- Intelligent content extraction using Trafilatura
|
||||
- BeautifulSoup fallback for complex pages
|
||||
- Configurable content length limits
|
||||
- Optional link extraction
|
||||
- Perfect for feeding webpage content to LLMs
|
||||
### Infrastructure Management
|
||||
- **Portainer Integration**: Stack and container management
|
||||
- **NPM Integration**: Nginx Proxy Manager domain and certificate management
|
||||
- **Service Control**: Start/stop services with container orchestration
|
||||
|
||||
### Home Automation (Housekeeping API)
|
||||
- **Device Control**: Turn on/off, toggle, and set brightness for smart devices
|
||||
- **Scene Activation**: Trigger Home Assistant scenes
|
||||
- **Script Execution**: Run Home Assistant scripts
|
||||
- **Automation Management**: Enable/disable automations
|
||||
- **State History**: Query device state changes over time
|
||||
- **Area Discovery**: List rooms and areas
|
||||
|
||||
### Utilities
|
||||
- **DNS Lookup**: Query DNS records (A, AAAA, MX, TXT, CNAME, NS, SOA, PTR)
|
||||
- **Health Checks**: Comprehensive service health monitoring
|
||||
- **AI Metrics Proxy**: Forward metrics requests to Core-AI service
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
src/
|
||||
├── config.py # Global application settings
|
||||
├── logging_config.py # Logging configuration
|
||||
├── base_schema.py # Base Pydantic models
|
||||
├── main.py # FastAPI application entry point
|
||||
└── web_scraper/ # Web scraper module
|
||||
├── __init__.py
|
||||
├── config.py # Module-specific settings
|
||||
├── schemas.py # Pydantic request/response models
|
||||
├── service.py # Business logic
|
||||
├── router.py # API routes
|
||||
└── exceptions.py # Custom exceptions
|
||||
├── config.py # Global application settings
|
||||
├── logging_config.py # Logging configuration
|
||||
├── base_controller.py # Base controller pattern
|
||||
├── main.py # FastAPI application entry point
|
||||
├── auth/
|
||||
│ └── oidc.py # OIDC authentication
|
||||
├── clients/
|
||||
│ ├── homeassistant_client.py # Home Assistant REST client
|
||||
│ ├── npm_client.py # Nginx Proxy Manager client
|
||||
│ ├── ollama_client.py # Ollama LLM client
|
||||
│ └── portainer_client.py # Portainer API client
|
||||
├── controllers/
|
||||
│ ├── ai_controller.py # AI metrics proxy
|
||||
│ ├── health_controller.py # Health endpoints
|
||||
│ ├── housekeeping_controller.py # Home automation endpoints
|
||||
│ ├── infrastructure_controller.py # Infrastructure management
|
||||
│ ├── static_controller.py # Static file serving
|
||||
│ └── tools_controller.py # DNS and utility tools
|
||||
└── dns/
|
||||
├── service.py # DNS lookup service
|
||||
└── exceptions.py # DNS-specific exceptions
|
||||
```
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Health
|
||||
- `GET /` - Service info and documentation links
|
||||
- `GET /health` - Basic health status
|
||||
- `GET /health/full` - Detailed component health
|
||||
- `GET /health/diagnostics` - Full diagnostic information
|
||||
|
||||
### Infrastructure (`/infrastructure`)
|
||||
- `GET /infrastructure/health` - Portainer/NPM connection status
|
||||
- `GET /infrastructure/services` - List all services (stacks)
|
||||
- `GET /infrastructure/services/{name}` - Get service details
|
||||
- `GET /infrastructure/services/{name}/status` - Service status
|
||||
- `POST /infrastructure/services/{name}/start` - Start service
|
||||
- `POST /infrastructure/services/{name}/stop` - Stop service
|
||||
- `GET /infrastructure/containers` - List containers
|
||||
- `GET /infrastructure/containers/{name}` - Container details
|
||||
- `GET /infrastructure/containers/{name}/logs` - Container logs
|
||||
- `GET /infrastructure/ports` - List exposed ports
|
||||
- `GET /infrastructure/domains` - List proxy domains
|
||||
- `GET /infrastructure/widget-data` - Dashboard widget data
|
||||
|
||||
### Housekeeping (`/housekeeping`)
|
||||
- `GET /housekeeping/health` - Home Assistant connection status
|
||||
- `GET /housekeeping/devices` - List controllable devices
|
||||
- `GET /housekeeping/devices/{entity_id}` - Device details
|
||||
- `POST /housekeeping/devices/{entity_id}/control` - Control device
|
||||
- `GET /housekeeping/scenes` - List scenes
|
||||
- `POST /housekeeping/scenes/{scene_id}/activate` - Activate scene
|
||||
- `GET /housekeeping/scripts` - List scripts
|
||||
- `POST /housekeeping/scripts/{script_id}/run` - Run script
|
||||
- `GET /housekeeping/automations` - List automations
|
||||
- `POST /housekeeping/automations/{automation_id}/toggle` - Toggle automation
|
||||
- `GET /housekeeping/history` - State history
|
||||
- `GET /housekeeping/areas` - List areas/rooms
|
||||
|
||||
### Tools (`/tools`)
|
||||
- `POST /tools/dns/lookup` - DNS record lookup
|
||||
|
||||
### AI (`/ai`)
|
||||
- `GET /ai/metrics` - Proxy to Core-AI metrics
|
||||
|
||||
## Development
|
||||
|
||||
### Requirements
|
||||
@@ -40,13 +103,30 @@ src/
|
||||
# Install dependencies
|
||||
pip install -r requirements.txt
|
||||
|
||||
# Copy credentials template
|
||||
cp src/credentials.example.py src/credentials.py
|
||||
# Edit src/credentials.py with your values
|
||||
|
||||
# Run locally
|
||||
uvicorn src.main:app --reload --host 0.0.0.0 --port 8083
|
||||
```
|
||||
|
||||
### Running Tests
|
||||
|
||||
```bash
|
||||
# Run all tests
|
||||
pytest
|
||||
|
||||
# Run with coverage
|
||||
pytest --cov=src --cov-report=term-missing
|
||||
|
||||
# Run specific test file
|
||||
pytest tests/test_housekeeping.py -v
|
||||
```
|
||||
|
||||
### Adding New Dependencies
|
||||
|
||||
**Important**: Dependencies use major version pinning (`~=`) for automatic patch updates while preventing breaking changes.
|
||||
Dependencies use major version pinning (`~=`) for automatic patch updates while preventing breaking changes.
|
||||
|
||||
1. Add package to `requirements.txt` with major version constraint:
|
||||
```
|
||||
@@ -58,14 +138,6 @@ uvicorn src.main:app --reload --host 0.0.0.0 --port 8083
|
||||
docker restart core-api
|
||||
```
|
||||
|
||||
The container automatically runs `pip install -r requirements.txt` on every boot, so new dependencies are installed immediately on restart.
|
||||
|
||||
**Version Pinning Best Practices**:
|
||||
- Use `~=` (compatible release) for most packages: `fastapi~=0.115.0`
|
||||
- Use `>=X,<Y` for complex constraints: `langchain-core>=0.3.17,<0.4.0`
|
||||
- Allows automatic security patches without breaking changes
|
||||
- Documented in PEP 440
|
||||
|
||||
### Docker Build
|
||||
|
||||
```bash
|
||||
@@ -88,130 +160,39 @@ docker run -p 8083:8083 core-code:latest
|
||||
|
||||
### Environment Variables
|
||||
|
||||
See `.env.example` for all available configuration options.
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `PORTAINER_URL` | Portainer API URL | `http://localhost:9000` |
|
||||
| `PORTAINER_API_KEY` | Portainer API key | - |
|
||||
| `NPM_URL` | Nginx Proxy Manager URL | `http://localhost:81` |
|
||||
| `NPM_EMAIL` | NPM admin email | - |
|
||||
| `NPM_PASSWORD` | NPM admin password | - |
|
||||
| `HOMEASSISTANT_URL` | Home Assistant URL | `http://localhost:8123` |
|
||||
| `HOMEASSISTANT_TOKEN` | HA long-lived access token | - |
|
||||
| `OLLAMA_URL` | Ollama API URL | `http://localhost:11434` |
|
||||
| `OIDC_ENABLED` | Enable OIDC auth | `false` |
|
||||
| `OIDC_ISSUER` | OIDC issuer URL | - |
|
||||
| `OIDC_AUDIENCE` | OIDC audience | - |
|
||||
|
||||
## API Documentation
|
||||
|
||||
Once deployed, access documentation at:
|
||||
- **Swagger UI**: http://192.168.86.149:8083/docs
|
||||
- **ReDoc**: http://192.168.86.149:8083/redoc
|
||||
- **OpenAPI Spec**: http://192.168.86.149:8083/openapi.json
|
||||
|
||||
## Integration with Open WebUI
|
||||
|
||||
### Method 1: Functions (OpenAPI Import)
|
||||
1. In Open WebUI, navigate to Functions
|
||||
2. Import from OpenAPI spec: `http://192.168.86.149:8083/openapi.json`
|
||||
3. Use functions directly in chat
|
||||
|
||||
### Method 2: Pipelines
|
||||
1. Create a pipeline that calls Core Code API endpoints
|
||||
2. Use as data source for LLM workflows
|
||||
|
||||
### Method 3: Direct API Calls
|
||||
```python
|
||||
import httpx
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
"http://192.168.86.149:8083/web-scraper/scrape",
|
||||
json={
|
||||
"url": "https://example.com",
|
||||
"extract_main_content": True
|
||||
}
|
||||
)
|
||||
data = response.json()
|
||||
```
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Web Scraper
|
||||
|
||||
**POST /web-scraper/scrape**
|
||||
|
||||
Scrape and extract content from a website.
|
||||
|
||||
Request:
|
||||
```json
|
||||
{
|
||||
"url": "https://example.com/article",
|
||||
"extract_main_content": true,
|
||||
"include_links": false,
|
||||
"max_length": 10000
|
||||
}
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"url": "https://example.com/article",
|
||||
"title": "Article Title",
|
||||
"content": "Extracted article content...",
|
||||
"extracted_at": "2025-11-12T19:30:00Z",
|
||||
"content_length": 5432,
|
||||
"links": null
|
||||
}
|
||||
```
|
||||
|
||||
## Logging
|
||||
|
||||
Logs are written to:
|
||||
- **Console**: stdout (captured by Docker)
|
||||
- **File**: `/app/logs/app.log` (persisted via volume mount)
|
||||
|
||||
Log format:
|
||||
```
|
||||
2025-11-12 19:30:00 | INFO | src.web_scraper.service:scrape_url:45 | Starting scrape for URL: https://example.com
|
||||
```
|
||||
- **Swagger UI**: http://localhost:8083/docs
|
||||
- **ReDoc**: http://localhost:8083/redoc
|
||||
- **OpenAPI Spec**: http://localhost:8083/openapi.json
|
||||
|
||||
## Health Checks
|
||||
|
||||
- **Endpoint**: `GET /health`
|
||||
- **Docker**: Automatic health checks configured
|
||||
- **Response**: `{"status": "healthy"}`
|
||||
- **Basic**: `GET /health` - Returns status and Ollama connection
|
||||
- **Full**: `GET /health/full` - Returns all component statuses (503 if unhealthy)
|
||||
- **Diagnostics**: `GET /health/diagnostics` - Detailed service information
|
||||
|
||||
## Security
|
||||
|
||||
- Runs as non-root user (uid 1000)
|
||||
- No authentication required (internal network only)
|
||||
- OIDC authentication support via Authentik
|
||||
- CORS configured for same-network access
|
||||
- Rate limiting: Not implemented (internal use only)
|
||||
|
||||
## Future Modules
|
||||
|
||||
The architecture supports adding new modules:
|
||||
- Data transformation functions
|
||||
- API integrations
|
||||
- File processing
|
||||
- Database queries
|
||||
|
||||
Each module follows the same structure:
|
||||
```
|
||||
src/
|
||||
└── module_name/
|
||||
├── config.py
|
||||
├── schemas.py
|
||||
├── service.py
|
||||
├── router.py
|
||||
└── exceptions.py
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Container won't start
|
||||
```bash
|
||||
docker logs core-code
|
||||
```
|
||||
|
||||
### API not responding
|
||||
```bash
|
||||
curl http://192.168.86.149:8083/health
|
||||
```
|
||||
|
||||
### Check OpenAPI spec
|
||||
```bash
|
||||
curl http://192.168.86.149:8083/openapi.json | jq
|
||||
```
|
||||
- Admin endpoints require authentication when OIDC enabled
|
||||
|
||||
## License
|
||||
|
||||
|
||||
@@ -1,446 +0,0 @@
|
||||
# Requested Core-API Services for Core-AI Infrastructure Tools
|
||||
|
||||
This document specifies the API endpoints needed by core-ai infrastructure tools. All requests from core-ai should go through core-api for centralized logging and access control.
|
||||
|
||||
## Context
|
||||
|
||||
The core-ai service is implementing 9 infrastructure tools in 3 logical clusters:
|
||||
1. **Container Lifecycle** (4 tools) - containers.py
|
||||
2. **Service Management** (3 tools) - services.py
|
||||
3. **Monitoring & Resources** (2 tools) - monitoring.py
|
||||
|
||||
These tools need corresponding core-api REST endpoints to perform operations via Portainer.
|
||||
|
||||
---
|
||||
|
||||
## Cluster 1: Container Lifecycle Management
|
||||
|
||||
### 1.1 List Containers
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/containers`
|
||||
|
||||
**Query Parameters:**
|
||||
- `status` (optional): Filter by status - "all", "running", "stopped", "paused" (default: "running")
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
[
|
||||
{
|
||||
"Id": "abc123...",
|
||||
"Names": ["/nginx"],
|
||||
"State": "running",
|
||||
"Status": "Up 3 days",
|
||||
"Image": "nginx:latest",
|
||||
"Ports": [
|
||||
{"PrivatePort": 80, "PublicPort": 8080, "Type": "tcp"},
|
||||
{"PrivatePort": 443, "PublicPort": 8443, "Type": "tcp"}
|
||||
],
|
||||
"StartedAt": "2024-12-01T10:00:00Z"
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Use `PortainerClient.list_containers(all_containers=True)` with Docker socket fallback
|
||||
- Filter results based on `status` query parameter
|
||||
- Return standard Docker API container list format
|
||||
|
||||
---
|
||||
|
||||
### 1.2 Manage Container
|
||||
|
||||
**Endpoint:** `POST /v1/infrastructure/containers/{container}/{action}`
|
||||
|
||||
**Path Parameters:**
|
||||
- `container`: Container name or ID (e.g., "nginx", "core-ai")
|
||||
- `action`: One of: "start", "stop", "restart", "pause", "unpause", "remove"
|
||||
|
||||
**Response (Success):**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"action": "restart",
|
||||
"container": "nginx",
|
||||
"message": "Container restarted successfully"
|
||||
}
|
||||
```
|
||||
|
||||
**Response (Error):**
|
||||
```json
|
||||
{
|
||||
"success": false,
|
||||
"error": "Container not found",
|
||||
"message": "Container 'nginx2' not found. Available containers: nginx, core-ai, ollama"
|
||||
}
|
||||
```
|
||||
|
||||
**Status Codes:**
|
||||
- `200` - Success
|
||||
- `304` - Not Modified (already in target state)
|
||||
- `404` - Container not found
|
||||
- `409` - Conflict (e.g., cannot remove running container)
|
||||
- `500` - Server error
|
||||
|
||||
**Implementation Notes:**
|
||||
- For "restart": call stop then start
|
||||
- For actions not yet in PortainerClient (pause, unpause, remove):
|
||||
- Call Portainer API directly: `/api/endpoints/{endpoint_id}/docker/containers/{container_id}/{action}`
|
||||
- Handle partial name matching (case-insensitive)
|
||||
- Return helpful error messages suggesting `docker_list_containers()` when not found
|
||||
|
||||
---
|
||||
|
||||
### 1.3 Inspect Container
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/containers/{container}`
|
||||
|
||||
**Path Parameters:**
|
||||
- `container`: Container name or ID
|
||||
|
||||
**Query Parameters:**
|
||||
- `details` (optional): Level of detail - "summary" (default), "full", "resources"
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"Id": "abc123...",
|
||||
"Name": "/nginx",
|
||||
"State": {
|
||||
"Status": "running",
|
||||
"Running": true,
|
||||
"StartedAt": "2024-12-01T10:00:00Z",
|
||||
"FinishedAt": "0001-01-01T00:00:00Z",
|
||||
"ExitCode": 0
|
||||
},
|
||||
"Config": {
|
||||
"Image": "nginx:latest",
|
||||
"Env": ["PATH=/usr/local/sbin:...", "NGINX_VERSION=1.25.0"],
|
||||
"Cmd": ["nginx", "-g", "daemon off;"]
|
||||
},
|
||||
"NetworkSettings": {
|
||||
"Ports": {
|
||||
"80/tcp": [{"HostIp": "0.0.0.0", "HostPort": "8080"}],
|
||||
"443/tcp": [{"HostIp": "0.0.0.0", "HostPort": "8443"}]
|
||||
},
|
||||
"Networks": {
|
||||
"bridge": {
|
||||
"IPAddress": "172.17.0.2",
|
||||
"Gateway": "172.17.0.1"
|
||||
}
|
||||
}
|
||||
},
|
||||
"HostConfig": {
|
||||
"Memory": 536870912,
|
||||
"NanoCpus": 1000000000,
|
||||
"RestartPolicy": {"Name": "unless-stopped"}
|
||||
},
|
||||
"Mounts": [
|
||||
{
|
||||
"Type": "bind",
|
||||
"Source": "/host/path",
|
||||
"Destination": "/container/path"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Use `PortainerClient.inspect_container(container)` which auto-detects endpoint and falls back to Docker socket
|
||||
- Return full Docker inspect response
|
||||
- The core-ai tool will handle formatting based on `details` level
|
||||
- Return 404 if container not found
|
||||
|
||||
---
|
||||
|
||||
### 1.4 Container Logs
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/containers/{container}/logs`
|
||||
|
||||
**Path Parameters:**
|
||||
- `container`: Container name or ID
|
||||
|
||||
**Query Parameters:**
|
||||
- `lines` (optional): Number of log lines (default: 50, max: 500)
|
||||
- `since` (optional): Time filter - "1h", "30m", or ISO timestamp
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"container": "nginx",
|
||||
"lines_requested": 50,
|
||||
"since": null,
|
||||
"logs": "2024-12-04T10:00:00.123Z Starting nginx...\n2024-12-04T10:00:01.456Z Ready to accept connections\n..."
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Access Docker API directly: `GET /v1.41/containers/{container}/logs`
|
||||
- Use Docker socket transport (httpx with uds)
|
||||
- Parameters: `stdout=true`, `stderr=true`, `tail={lines}`, `timestamps=true`
|
||||
- If `since` provided: add `since={unix_timestamp}` parameter
|
||||
- Strip Docker stream headers (8-byte binary prefix per line)
|
||||
- Return plain text logs with timestamps
|
||||
- Return 404 if container not found
|
||||
|
||||
---
|
||||
|
||||
## Cluster 2: Service Management
|
||||
|
||||
### 2.1 List Services
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/services` (already exists, may need enhancement)
|
||||
|
||||
**Query Parameters:**
|
||||
- `stack` (optional): Filter by stack name
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
[
|
||||
{
|
||||
"name": "portainer",
|
||||
"stack_id": 1,
|
||||
"status": "active",
|
||||
"containers_running": 3,
|
||||
"containers_total": 3,
|
||||
"ports": [9000, 8000],
|
||||
"domains": ["portainer.example.com"]
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Enhance existing `/infrastructure/services` endpoint if needed
|
||||
- Ensure it returns stack/service information from Portainer
|
||||
- Include container counts (running/total)
|
||||
|
||||
---
|
||||
|
||||
### 2.2 Manage Service
|
||||
|
||||
**Endpoint:** `POST /v1/infrastructure/services/{service}/{action}`
|
||||
|
||||
**Path Parameters:**
|
||||
- `service`: Service/stack name
|
||||
- `action`: One of: "start", "stop", "restart", "scale"
|
||||
|
||||
**Request Body (for scale action):**
|
||||
```json
|
||||
{
|
||||
"replicas": 3
|
||||
}
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"action": "restart",
|
||||
"service": "web",
|
||||
"message": "Service restarted successfully"
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- For "start"/"stop": Use Portainer stack start/stop API
|
||||
- For "restart": Stop then start the stack
|
||||
- For "scale": Update stack with new replica count
|
||||
- This may require updating stack compose file
|
||||
|
||||
---
|
||||
|
||||
### 2.3 Service Status
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/services/{service}/status`
|
||||
|
||||
**Path Parameters:**
|
||||
- `service`: Service/stack name
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"name": "web",
|
||||
"status": "active",
|
||||
"stack_id": 5,
|
||||
"containers": [
|
||||
{
|
||||
"name": "web_app_1",
|
||||
"status": "running",
|
||||
"health": "healthy",
|
||||
"uptime": "2 days"
|
||||
}
|
||||
],
|
||||
"replica_status": "3/3 running",
|
||||
"resources": {
|
||||
"memory_total": "1.2 GB",
|
||||
"cpu_usage": "15%"
|
||||
},
|
||||
"recent_events": [
|
||||
{"time": "2024-12-04T09:00:00Z", "action": "container_start", "container": "web_app_3"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Get stack details from Portainer
|
||||
- Get individual container statuses
|
||||
- Calculate aggregate resource usage
|
||||
- May require querying Docker events API for recent events
|
||||
|
||||
---
|
||||
|
||||
## Cluster 3: Monitoring & Resources
|
||||
|
||||
### 3.1 System Resources
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/resources/system`
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"cpu": {
|
||||
"cores": 8,
|
||||
"usage_percent": 45.2,
|
||||
"load_average": [2.5, 2.3, 2.1]
|
||||
},
|
||||
"memory": {
|
||||
"total_bytes": 16777216000,
|
||||
"used_bytes": 8388608000,
|
||||
"available_bytes": 8388608000,
|
||||
"usage_percent": 50.0
|
||||
},
|
||||
"disk": {
|
||||
"total_bytes": 500000000000,
|
||||
"used_bytes": 250000000000,
|
||||
"available_bytes": 250000000000,
|
||||
"usage_percent": 50.0
|
||||
},
|
||||
"network": {
|
||||
"interfaces": {
|
||||
"eth0": {
|
||||
"rx_bytes": 1000000000,
|
||||
"tx_bytes": 500000000
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Use Docker system info API: `GET /v1.41/system/df`
|
||||
- May also use `GET /v1.41/info` for system-wide stats
|
||||
- Calculate percentages and format nicely
|
||||
- Include load averages from system stats
|
||||
|
||||
---
|
||||
|
||||
### 3.2 Container Resources
|
||||
|
||||
**Endpoint:** `GET /v1/infrastructure/resources/containers`
|
||||
|
||||
**Query Parameters:**
|
||||
- `container` (optional): Specific container name/ID (if omitted, return all)
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
[
|
||||
{
|
||||
"name": "nginx",
|
||||
"cpu_percent": 5.2,
|
||||
"memory_usage_bytes": 45000000,
|
||||
"memory_limit_bytes": 100000000,
|
||||
"memory_percent": 45.0,
|
||||
"network_rx_bytes": 50000000,
|
||||
"network_tx_bytes": 25000000,
|
||||
"block_read_bytes": 10000000,
|
||||
"block_write_bytes": 5000000
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
**Implementation Notes:**
|
||||
- Use Docker stats API: `GET /v1.41/containers/{id}/stats?stream=false`
|
||||
- If `container` param provided: return single container stats
|
||||
- If omitted: return stats for all running containers
|
||||
- Calculate percentages where applicable
|
||||
- Stats API returns real-time metrics (one-time snapshot, not streaming)
|
||||
|
||||
---
|
||||
|
||||
## Implementation Priority
|
||||
|
||||
**Phase 1 (Needed immediately for core-ai):**
|
||||
1. `GET /v1/infrastructure/containers` - List containers
|
||||
2. `POST /v1/infrastructure/containers/{container}/{action}` - Manage containers
|
||||
3. `GET /v1/infrastructure/containers/{container}` - Inspect container
|
||||
4. `GET /v1/infrastructure/containers/{container}/logs` - Container logs
|
||||
|
||||
**Phase 2 (Needed for full infrastructure tools):**
|
||||
5. `POST /v1/infrastructure/services/{service}/{action}` - Manage services
|
||||
6. `GET /v1/infrastructure/services/{service}/status` - Service status
|
||||
7. `GET /v1/infrastructure/resources/system` - System resources
|
||||
8. `GET /v1/infrastructure/resources/containers` - Container resources
|
||||
|
||||
---
|
||||
|
||||
## Security & Access Control
|
||||
|
||||
All endpoints should:
|
||||
- Log all requests (especially write operations)
|
||||
- Support OIDC authentication when enabled
|
||||
- Require admin privileges for destructive operations (remove, scale)
|
||||
- Rate limit to prevent abuse
|
||||
- Validate input parameters
|
||||
- Return sanitized errors (no sensitive data in error messages)
|
||||
|
||||
---
|
||||
|
||||
## Error Handling
|
||||
|
||||
Standard error response format:
|
||||
```json
|
||||
{
|
||||
"error": "ContainerNotFound",
|
||||
"message": "Container 'nginx2' not found",
|
||||
"details": {
|
||||
"container": "nginx2",
|
||||
"available_containers": ["nginx", "core-ai", "ollama"]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Common error codes:
|
||||
- `400` - Bad Request (invalid parameters)
|
||||
- `404` - Not Found (container/service doesn't exist)
|
||||
- `409` - Conflict (invalid state transition)
|
||||
- `500` - Internal Server Error (Portainer/Docker API failed)
|
||||
- `503` - Service Unavailable (Portainer/Docker not accessible)
|
||||
|
||||
---
|
||||
|
||||
## Testing
|
||||
|
||||
Each endpoint should have:
|
||||
- Unit tests (mock Portainer client)
|
||||
- Integration tests (real Portainer/Docker)
|
||||
- Error case tests (not found, permission denied, etc.)
|
||||
- Performance tests (ensure response times < 2s)
|
||||
|
||||
---
|
||||
|
||||
## Questions / Decisions Needed
|
||||
|
||||
1. **Authentication**: Should container management require admin role, or allow read-only for all users?
|
||||
2. **Rate Limiting**: What limits should be applied to prevent abuse?
|
||||
3. **Caching**: Should container lists be cached? (TTL: 5s?)
|
||||
4. **Async**: Should heavy operations (like logs) be async with job IDs?
|
||||
5. **Webhooks**: Should operations emit events for monitoring?
|
||||
|
||||
---
|
||||
|
||||
## Notes
|
||||
|
||||
- All endpoints follow RESTful conventions
|
||||
- Use existing PortainerClient methods where available
|
||||
- Fall back to Docker socket when Portainer doesn't have data
|
||||
- Log all operations with timestamps, user, and outcome
|
||||
- Consider adding `/v1/infrastructure/containers/search` for fuzzy name matching
|
||||
+68
@@ -0,0 +1,68 @@
|
||||
# Alembic Configuration for Core-API
|
||||
#
|
||||
# Database migrations using async SQLAlchemy
|
||||
|
||||
[alembic]
|
||||
# Path to migration scripts
|
||||
script_location = alembic
|
||||
|
||||
# Template for migration file names
|
||||
file_template = %%(year)d%%(month).2d%%(day).2d_%%(hour).2d%%(minute).2d_%%(rev)s_%%(slug)s
|
||||
|
||||
# Prepend sys.path with the project root
|
||||
prepend_sys_path = .
|
||||
|
||||
# Timezone for revision creation date
|
||||
timezone = UTC
|
||||
|
||||
# Max length of revision identifiers
|
||||
truncate_slug_length = 40
|
||||
|
||||
# Set to 'true' to run in offline mode
|
||||
revision_environment = false
|
||||
|
||||
# Set to 'true' for sqlalchemy.url to be from the environment
|
||||
# We use env.py to get the URL from config.py instead
|
||||
sqlalchemy.url =
|
||||
|
||||
[post_write_hooks]
|
||||
# Black formatting on generated migration files
|
||||
# hooks = black
|
||||
# black.type = console_scripts
|
||||
# black.entrypoint = black
|
||||
# black.options = -l 100 REVISION_SCRIPT_FILENAME
|
||||
|
||||
# Logging configuration
|
||||
[loggers]
|
||||
keys = root,sqlalchemy,alembic
|
||||
|
||||
[handlers]
|
||||
keys = console
|
||||
|
||||
[formatters]
|
||||
keys = generic
|
||||
|
||||
[logger_root]
|
||||
level = WARN
|
||||
handlers = console
|
||||
qualname =
|
||||
|
||||
[logger_sqlalchemy]
|
||||
level = WARN
|
||||
handlers =
|
||||
qualname = sqlalchemy.engine
|
||||
|
||||
[logger_alembic]
|
||||
level = INFO
|
||||
handlers =
|
||||
qualname = alembic
|
||||
|
||||
[handler_console]
|
||||
class = StreamHandler
|
||||
args = (sys.stderr,)
|
||||
level = NOTSET
|
||||
formatter = generic
|
||||
|
||||
[formatter_generic]
|
||||
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||
datefmt = %H:%M:%S
|
||||
@@ -0,0 +1,21 @@
|
||||
Alembic Migrations for Core-API
|
||||
|
||||
This directory contains database migrations managed by Alembic.
|
||||
|
||||
Commands:
|
||||
# Generate a new migration (after changing models)
|
||||
alembic revision --autogenerate -m "description"
|
||||
|
||||
# Apply all pending migrations
|
||||
alembic upgrade head
|
||||
|
||||
# Rollback last migration
|
||||
alembic downgrade -1
|
||||
|
||||
# View migration history
|
||||
alembic history
|
||||
|
||||
# View current revision
|
||||
alembic current
|
||||
|
||||
See https://alembic.sqlalchemy.org for more documentation.
|
||||
@@ -0,0 +1,99 @@
|
||||
"""
|
||||
Alembic Environment Configuration
|
||||
|
||||
Async migration environment for SQLAlchemy 2.0 with asyncpg.
|
||||
"""
|
||||
import asyncio
|
||||
from logging.config import fileConfig
|
||||
|
||||
from sqlalchemy import pool
|
||||
from sqlalchemy.engine import Connection
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from alembic import context
|
||||
|
||||
# Import our models and config
|
||||
from src.shared.config import get_settings
|
||||
from src.shared.database import Base
|
||||
|
||||
# Import all models to ensure they're registered with Base.metadata
|
||||
from src.domains.auth.models import User, Role, UserRole, UserPreferences, ApiKey # noqa: F401
|
||||
from src.domains.dashboard.models import QuickLink, DashboardWidget # noqa: F401
|
||||
|
||||
# Alembic Config object
|
||||
config = context.config
|
||||
|
||||
# Interpret the config file for Python logging
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
# Model metadata for autogenerate support
|
||||
target_metadata = Base.metadata
|
||||
|
||||
# Get database URL from our settings
|
||||
settings = get_settings()
|
||||
db_url = settings.database_url
|
||||
if db_url.startswith("postgresql://"):
|
||||
db_url = db_url.replace("postgresql://", "postgresql+asyncpg://", 1)
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
"""
|
||||
Run migrations in 'offline' mode.
|
||||
|
||||
Generates SQL script without connecting to the database.
|
||||
"""
|
||||
context.configure(
|
||||
url=db_url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def do_run_migrations(connection: Connection) -> None:
|
||||
"""
|
||||
Run migrations with the given connection.
|
||||
"""
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
async def run_async_migrations() -> None:
|
||||
"""
|
||||
Run migrations in 'online' mode with async engine.
|
||||
"""
|
||||
configuration = config.get_section(config.config_ini_section) or {}
|
||||
configuration["sqlalchemy.url"] = db_url
|
||||
|
||||
connectable = async_engine_from_config(
|
||||
configuration,
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
|
||||
async with connectable.connect() as connection:
|
||||
await connection.run_sync(do_run_migrations)
|
||||
|
||||
await connectable.dispose()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
"""
|
||||
Run migrations in 'online' mode.
|
||||
"""
|
||||
asyncio.run(run_async_migrations())
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
@@ -0,0 +1,25 @@
|
||||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
${imports if imports else ""}
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = ${repr(up_revision)}
|
||||
down_revision: Union[str, None] = ${repr(down_revision)}
|
||||
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
|
||||
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
${downgrades if downgrades else "pass"}
|
||||
@@ -0,0 +1,138 @@
|
||||
"""Create auth tables
|
||||
|
||||
Revision ID: 001
|
||||
Revises:
|
||||
Create Date: 2026-01-01
|
||||
|
||||
Creates the initial authentication and authorization tables:
|
||||
- users: User accounts synced from Authentik
|
||||
- roles: Domain-scoped permission roles
|
||||
- user_roles: User-Role association table
|
||||
- user_preferences: User settings and preferences
|
||||
- api_keys: API key authentication
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "001"
|
||||
down_revision: Union[str, None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Users table
|
||||
op.create_table(
|
||||
"users",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
|
||||
sa.Column("authentik_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("email", sa.String(255), nullable=False),
|
||||
sa.Column("name", sa.String(255), nullable=False),
|
||||
sa.Column("avatar_url", sa.String(500), nullable=True),
|
||||
sa.Column("api_keys_enabled", sa.Boolean(), nullable=False, server_default="true"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("last_login", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
op.create_index("ix_users_authentik_id", "users", ["authentik_id"], unique=True)
|
||||
op.create_index("ix_users_email", "users", ["email"], unique=True)
|
||||
|
||||
# Roles table
|
||||
op.create_table(
|
||||
"roles",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
|
||||
sa.Column("name", sa.String(100), nullable=False, comment="Role name in format domain:action"),
|
||||
sa.Column("domain", sa.String(50), nullable=False, comment="Permission domain"),
|
||||
sa.Column("action", sa.String(20), nullable=False, comment="Permission action"),
|
||||
sa.Column("authentik_group", sa.String(255), nullable=True, comment="Corresponding Authentik group name"),
|
||||
)
|
||||
op.create_index("ix_roles_name", "roles", ["name"], unique=True)
|
||||
op.create_index("ix_roles_domain", "roles", ["domain"])
|
||||
op.create_index("ix_roles_authentik_group", "roles", ["authentik_group"], unique=True)
|
||||
|
||||
# User-Role association table
|
||||
op.create_table(
|
||||
"user_roles",
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("users.id", ondelete="CASCADE"), primary_key=True),
|
||||
sa.Column("role_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("roles.id", ondelete="CASCADE"), primary_key=True),
|
||||
)
|
||||
|
||||
# User preferences table
|
||||
op.create_table(
|
||||
"user_preferences",
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("users.id", ondelete="CASCADE"), primary_key=True),
|
||||
sa.Column("theme", sa.String(20), nullable=False, server_default="system", comment="Theme preference"),
|
||||
sa.Column("default_room", sa.String(50), nullable=False, server_default="front-hall", comment="Default room"),
|
||||
sa.Column("preferences_json", postgresql.JSONB(), nullable=False, server_default="{}"),
|
||||
)
|
||||
|
||||
# API keys table
|
||||
op.create_table(
|
||||
"api_keys",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False),
|
||||
sa.Column("name", sa.String(100), nullable=False, comment="Human-readable key name"),
|
||||
sa.Column("key_hash", sa.String(255), nullable=False, comment="SHA-256 hash of the API key"),
|
||||
sa.Column("key_prefix", sa.String(8), nullable=False, comment="First 8 chars for identification"),
|
||||
sa.Column("scopes", postgresql.ARRAY(sa.String()), nullable=True, comment="Optional scope restriction"),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("last_used_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
)
|
||||
op.create_index("ix_api_keys_user_id", "api_keys", ["user_id"])
|
||||
|
||||
# Seed initial roles (domain:action combinations)
|
||||
roles_table = sa.table(
|
||||
"roles",
|
||||
sa.column("id", postgresql.UUID),
|
||||
sa.column("name", sa.String),
|
||||
sa.column("domain", sa.String),
|
||||
sa.column("action", sa.String),
|
||||
sa.column("authentik_group", sa.String),
|
||||
)
|
||||
|
||||
domains = [
|
||||
"control-room",
|
||||
"library",
|
||||
"media",
|
||||
"ai",
|
||||
"housekeeper",
|
||||
"developer",
|
||||
"documents",
|
||||
"gaming",
|
||||
"admin",
|
||||
]
|
||||
actions = ["viewer", "user", "editor", "admin"]
|
||||
|
||||
roles_data = []
|
||||
for domain in domains:
|
||||
for action in actions:
|
||||
role_name = f"{domain}:{action}"
|
||||
authentik_group = f"tatlock-{domain}-{action}"
|
||||
roles_data.append({
|
||||
"id": sa.text("gen_random_uuid()"),
|
||||
"name": role_name,
|
||||
"domain": domain,
|
||||
"action": action,
|
||||
"authentik_group": authentik_group,
|
||||
})
|
||||
|
||||
# Insert roles using raw SQL for UUID generation
|
||||
for role in roles_data:
|
||||
op.execute(
|
||||
f"""
|
||||
INSERT INTO roles (id, name, domain, action, authentik_group)
|
||||
VALUES (gen_random_uuid(), '{role["name"]}', '{role["domain"]}', '{role["action"]}', '{role["authentik_group"]}')
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("api_keys")
|
||||
op.drop_table("user_preferences")
|
||||
op.drop_table("user_roles")
|
||||
op.drop_table("roles")
|
||||
op.drop_table("users")
|
||||
@@ -0,0 +1,40 @@
|
||||
"""Create groups table
|
||||
|
||||
Revision ID: 002
|
||||
Revises: 001
|
||||
Create Date: 2026-01-01
|
||||
|
||||
Creates the groups table for syncing Authentik groups.
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "002"
|
||||
down_revision: Union[str, None] = "001"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Groups table
|
||||
op.create_table(
|
||||
"groups",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
|
||||
sa.Column("authentik_id", postgresql.UUID(as_uuid=True), nullable=False, comment="Authentik group UUID"),
|
||||
sa.Column("name", sa.String(255), nullable=False),
|
||||
sa.Column("is_superuser", sa.Boolean(), nullable=False, server_default="false", comment="Whether members have superuser privileges"),
|
||||
sa.Column("parent_name", sa.String(255), nullable=True, comment="Parent group name for hierarchy"),
|
||||
sa.Column("member_count", sa.Integer(), nullable=False, server_default="0", comment="Number of users in this group"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("synced_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False, comment="Last sync from Authentik"),
|
||||
)
|
||||
op.create_index("ix_groups_authentik_id", "groups", ["authentik_id"], unique=True)
|
||||
op.create_index("ix_groups_name", "groups", ["name"], unique=True)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("groups")
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Create dashboard tables
|
||||
|
||||
Revision ID: 003
|
||||
Revises: 002
|
||||
Create Date: 2026-01-03
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '003'
|
||||
down_revision: Union[str, None] = '002'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create quick_links and dashboard_widgets tables."""
|
||||
# Create quick_links table
|
||||
op.create_table(
|
||||
'quick_links',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('title', sa.String(length=100), nullable=False),
|
||||
sa.Column('url', sa.String(length=500), nullable=False),
|
||||
sa.Column('icon', sa.String(length=100), nullable=True),
|
||||
sa.Column('description', sa.String(length=255), nullable=True),
|
||||
sa.Column('category', sa.String(length=50), nullable=True),
|
||||
sa.Column('user_id', sa.String(length=255), nullable=True),
|
||||
sa.Column('position', sa.Integer(), nullable=True, default=0),
|
||||
sa.Column('is_visible', sa.Boolean(), nullable=True, default=True),
|
||||
sa.Column('color', sa.String(length=20), nullable=True),
|
||||
sa.Column('background_color', sa.String(length=20), nullable=True),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=True),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_quick_links_id'), 'quick_links', ['id'], unique=False)
|
||||
op.create_index(op.f('ix_quick_links_user_id'), 'quick_links', ['user_id'], unique=False)
|
||||
|
||||
# Create dashboard_widgets table
|
||||
op.create_table(
|
||||
'dashboard_widgets',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('widget_type', sa.String(length=50), nullable=False),
|
||||
sa.Column('user_id', sa.String(length=255), nullable=True),
|
||||
sa.Column('position_x', sa.Integer(), nullable=True, default=0),
|
||||
sa.Column('position_y', sa.Integer(), nullable=True, default=0),
|
||||
sa.Column('width', sa.Integer(), nullable=True, default=1),
|
||||
sa.Column('height', sa.Integer(), nullable=True, default=1),
|
||||
sa.Column('config', sa.Text(), nullable=True),
|
||||
sa.Column('is_visible', sa.Boolean(), nullable=True, default=True),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=True),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_dashboard_widgets_id'), 'dashboard_widgets', ['id'], unique=False)
|
||||
op.create_index(op.f('ix_dashboard_widgets_user_id'), 'dashboard_widgets', ['user_id'], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop dashboard tables."""
|
||||
op.drop_index(op.f('ix_dashboard_widgets_user_id'), table_name='dashboard_widgets')
|
||||
op.drop_index(op.f('ix_dashboard_widgets_id'), table_name='dashboard_widgets')
|
||||
op.drop_table('dashboard_widgets')
|
||||
|
||||
op.drop_index(op.f('ix_quick_links_user_id'), table_name='quick_links')
|
||||
op.drop_index(op.f('ix_quick_links_id'), table_name='quick_links')
|
||||
op.drop_table('quick_links')
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add_link_type_to_quick_links
|
||||
|
||||
Revision ID: f0349c95aa5d
|
||||
Revises: 003
|
||||
Create Date: 2026-01-03 13:14:46.911770+00:00
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'f0349c95aa5d'
|
||||
down_revision: Union[str, None] = '003'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column('quick_links', sa.Column('link_type', sa.String(length=20), nullable=True, server_default='iframe'))
|
||||
# Update existing rows to have the default value
|
||||
op.execute("UPDATE quick_links SET link_type = 'iframe' WHERE link_type IS NULL")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column('quick_links', 'link_type')
|
||||
@@ -0,0 +1,141 @@
|
||||
"""Add group_roles mapping and update role schema
|
||||
|
||||
Revision ID: 004
|
||||
Revises: f0349c95aa5d
|
||||
Create Date: 2026-01-03
|
||||
|
||||
Changes:
|
||||
- Add category column to roles (default 'general')
|
||||
- Drop authentik_group column from roles (decoupled architecture)
|
||||
- Create user_groups association table
|
||||
- Create group_roles association table
|
||||
- Update role names from domain:action to domain.general:action
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "004"
|
||||
down_revision: Union[str, None] = "f0349c95aa5d"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add category column to roles
|
||||
op.add_column(
|
||||
"roles",
|
||||
sa.Column(
|
||||
"category",
|
||||
sa.String(50),
|
||||
nullable=False,
|
||||
server_default="general",
|
||||
comment="Permission category within domain (general for full access, or specific tool)",
|
||||
),
|
||||
)
|
||||
|
||||
# Update role names from domain:action to domain.general:action
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE roles
|
||||
SET name = REPLACE(name, ':', '.general:')
|
||||
WHERE name NOT LIKE '%.%:%'
|
||||
"""
|
||||
)
|
||||
|
||||
# Update the comment on the name column
|
||||
op.alter_column(
|
||||
"roles",
|
||||
"name",
|
||||
comment="Role name in format domain.category:action (e.g., control-room.general:admin)",
|
||||
)
|
||||
|
||||
# Drop the authentik_group unique index first
|
||||
op.drop_index("ix_roles_authentik_group", table_name="roles")
|
||||
|
||||
# Drop authentik_group column (no longer needed with group_roles mapping)
|
||||
op.drop_column("roles", "authentik_group")
|
||||
|
||||
# Create user_groups association table
|
||||
op.create_table(
|
||||
"user_groups",
|
||||
sa.Column(
|
||||
"user_id",
|
||||
postgresql.UUID(as_uuid=True),
|
||||
sa.ForeignKey("users.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
),
|
||||
sa.Column(
|
||||
"group_id",
|
||||
postgresql.UUID(as_uuid=True),
|
||||
sa.ForeignKey("groups.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
),
|
||||
)
|
||||
|
||||
# Create group_roles association table
|
||||
op.create_table(
|
||||
"group_roles",
|
||||
sa.Column(
|
||||
"group_id",
|
||||
postgresql.UUID(as_uuid=True),
|
||||
sa.ForeignKey("groups.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
),
|
||||
sa.Column(
|
||||
"role_id",
|
||||
postgresql.UUID(as_uuid=True),
|
||||
sa.ForeignKey("roles.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop association tables
|
||||
op.drop_table("group_roles")
|
||||
op.drop_table("user_groups")
|
||||
|
||||
# Add back authentik_group column
|
||||
op.add_column(
|
||||
"roles",
|
||||
sa.Column(
|
||||
"authentik_group",
|
||||
sa.String(255),
|
||||
nullable=True,
|
||||
comment="Corresponding Authentik group name",
|
||||
),
|
||||
)
|
||||
|
||||
# Restore authentik_group values from role names
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE roles
|
||||
SET authentik_group = 'tatlock-' || REPLACE(REPLACE(name, '.general:', '-'), ':', '-')
|
||||
"""
|
||||
)
|
||||
|
||||
# Recreate the unique index
|
||||
op.create_index("ix_roles_authentik_group", "roles", ["authentik_group"], unique=True)
|
||||
|
||||
# Revert role names from domain.general:action to domain:action
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE roles
|
||||
SET name = REPLACE(name, '.general:', ':')
|
||||
WHERE name LIKE '%.general:%'
|
||||
"""
|
||||
)
|
||||
|
||||
# Update the comment on the name column
|
||||
op.alter_column(
|
||||
"roles",
|
||||
"name",
|
||||
comment="Role name in format domain:action",
|
||||
)
|
||||
|
||||
# Drop category column
|
||||
op.drop_column("roles", "category")
|
||||
@@ -0,0 +1,8 @@
|
||||
# Development and testing dependencies
|
||||
-r requirements.txt
|
||||
|
||||
# Testing
|
||||
pytest>=9.0.0
|
||||
pytest-asyncio>=0.24.0
|
||||
pytest-cov>=6.0.0
|
||||
httpx>=0.28.0 # For TestClient
|
||||
+7
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "core-api"
|
||||
version = "1.0.0"
|
||||
version = "1.9.0"
|
||||
description = "Core Code API - Infrastructure management and tools API"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
@@ -8,3 +8,9 @@ license = {text = "MIT"}
|
||||
|
||||
[tool.setuptools]
|
||||
packages = ["src"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
pythonpath = ["."]
|
||||
asyncio_mode = "auto"
|
||||
addopts = "-v"
|
||||
|
||||
+15
-8
@@ -1,12 +1,13 @@
|
||||
# FastAPI and ASGI server
|
||||
fastapi~=0.115.0
|
||||
uvicorn[standard]>=0.34.0 # Updated for google-adk compatibility
|
||||
pydantic>=2.11.1,<3.0.0 # Required for google-cloud-aiplatform[agent-engines]
|
||||
fastapi>=0.115.0
|
||||
starlette>=0.49.1 # CVE-2025-54121, CVE-2025-62727
|
||||
uvicorn[standard]>=0.34.0
|
||||
pydantic>=2.11.1,<3.0.0
|
||||
pydantic-settings>=2.10.1
|
||||
|
||||
# HTTP client
|
||||
httpx>=0.28.0 # Required for google-adk
|
||||
python-socketio[asyncio_client]~=5.11.0
|
||||
httpx>=0.28.0
|
||||
python-socketio[asyncio_client]>=5.14.0 # CVE-2025-61765
|
||||
|
||||
# Web scraping
|
||||
beautifulsoup4~=4.12.0
|
||||
@@ -20,8 +21,14 @@ python-dotenv~=1.0.0
|
||||
python-json-logger~=2.0.0
|
||||
pytz~=2024.1
|
||||
dnspython~=2.7.0
|
||||
psutil~=6.1.0
|
||||
|
||||
# Authentication & Security
|
||||
PyJWT[crypto]~=2.9.0
|
||||
python-jose[cryptography]~=3.3.0
|
||||
cryptography~=43.0.0
|
||||
PyJWT[crypto]>=2.9.0
|
||||
python-jose[cryptography]>=3.4.0 # CVE PYSEC-2024-232, PYSEC-2024-233
|
||||
cryptography>=44.0.1 # CVE-2024-12797
|
||||
|
||||
# Database
|
||||
sqlalchemy[asyncio]~=2.0.0
|
||||
asyncpg>=0.30.0
|
||||
alembic~=1.13.0
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
"""
|
||||
Core Code API - OpenAPI-compatible functions for Open WebUI
|
||||
"""
|
||||
from src.shared.config import __version__
|
||||
|
||||
__version__ = "1.0.0"
|
||||
__author__ = "Core Code Team"
|
||||
|
||||
@@ -3,3 +3,25 @@ Authentication module for core-api
|
||||
|
||||
Provides OIDC/OAuth2 authentication via Authentik
|
||||
"""
|
||||
from src.auth.oidc import (
|
||||
get_current_user,
|
||||
get_admin_user,
|
||||
get_optional_user,
|
||||
get_forward_auth_user,
|
||||
get_forward_auth_admin,
|
||||
oidc_config,
|
||||
)
|
||||
from src.auth.service import AuthService, get_auth_service
|
||||
from src.auth.controller import auth_controller
|
||||
|
||||
__all__ = [
|
||||
"get_current_user",
|
||||
"get_admin_user",
|
||||
"get_optional_user",
|
||||
"get_forward_auth_user",
|
||||
"get_forward_auth_admin",
|
||||
"oidc_config",
|
||||
"AuthService",
|
||||
"get_auth_service",
|
||||
"auth_controller",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,309 @@
|
||||
"""
|
||||
Authentication Controller
|
||||
|
||||
Provides authentication endpoints for OIDC token sync and user management.
|
||||
"""
|
||||
from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.controllers.base import BaseController
|
||||
from src.logging_config import get_logger
|
||||
from src.db import get_async_session
|
||||
from src.auth.schemas import AuthSyncRequest, AuthSyncResponse, UsersListResponse, BulkSyncResultSchema, GroupsListResponse
|
||||
from src.auth.service import AuthService
|
||||
from src.auth.oidc import get_forward_auth_user
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AuthController(BaseController):
|
||||
"""
|
||||
Controller for authentication operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Token synchronization (login)
|
||||
- User profile retrieval
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/auth", tags=["Authentication"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.post(
|
||||
"/sync",
|
||||
summary="Sync user from OIDC token",
|
||||
response_model=AuthSyncResponse,
|
||||
responses={
|
||||
200: {"description": "User synced successfully"},
|
||||
401: {"description": "Invalid or expired token"},
|
||||
503: {"description": "Authentication service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_user(
|
||||
request: AuthSyncRequest,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> AuthSyncResponse:
|
||||
"""
|
||||
Synchronize user from OIDC access token
|
||||
|
||||
This endpoint should be called after the client obtains an access token
|
||||
from Authentik. It:
|
||||
1. Validates the token via Authentik's userinfo endpoint
|
||||
2. Creates or updates the user in the database
|
||||
3. Syncs roles from Authentik groups
|
||||
4. Returns the user profile with roles and preferences
|
||||
|
||||
The client should store the returned user info for local use.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
# Validate token with Authentik
|
||||
token_info = await service.validate_token(request.access_token)
|
||||
except ValueError as e:
|
||||
logger.warning(f"Token validation failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
# Sync user to database
|
||||
user, is_new = await service.sync_user(token_info)
|
||||
|
||||
# Sync roles from groups
|
||||
roles = await service.sync_roles(user, token_info.groups)
|
||||
|
||||
# Commit the transaction
|
||||
await session.commit()
|
||||
|
||||
# Refresh to get relationships
|
||||
await session.refresh(user, ["preferences"])
|
||||
|
||||
# Build response
|
||||
return AuthSyncResponse(
|
||||
user=service.user_to_schema(user),
|
||||
roles=service.roles_to_schema(roles),
|
||||
preferences=service.preferences_to_schema(user.preferences),
|
||||
is_new_user=is_new,
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/users",
|
||||
summary="List all users",
|
||||
response_model=UsersListResponse,
|
||||
responses={
|
||||
200: {"description": "List of users"},
|
||||
},
|
||||
)
|
||||
async def list_users(
|
||||
search: Optional[str] = Query(None, description="Search by name or email"),
|
||||
offset: int = Query(0, ge=0, description="Number of records to skip"),
|
||||
limit: int = Query(50, ge=1, le=100, description="Maximum records to return"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> UsersListResponse:
|
||||
"""
|
||||
List all users who have logged in via Authentik
|
||||
|
||||
Returns paginated list of users with their roles.
|
||||
Supports search filtering by name or email.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
items, total = await service.list_users(
|
||||
search=search,
|
||||
offset=offset,
|
||||
limit=limit,
|
||||
)
|
||||
return UsersListResponse(items=items, total=total)
|
||||
|
||||
@router.post(
|
||||
"/users/sync-from-authentik",
|
||||
summary="Bulk sync users from Authentik",
|
||||
response_model=BulkSyncResultSchema,
|
||||
responses={
|
||||
200: {"description": "Sync completed"},
|
||||
401: {"description": "Authentik API token invalid"},
|
||||
503: {"description": "Authentik service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_users_from_authentik(
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all users from Authentik and sync to local database
|
||||
|
||||
This endpoint uses the Authentik admin API to fetch all users
|
||||
and create/update them in the local database. Requires
|
||||
AUTHENTIK_CORE_API_TOKEN to be configured.
|
||||
|
||||
Use this to initially populate users or to re-sync after
|
||||
changes in Authentik.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
result = await service.bulk_sync_from_authentik()
|
||||
logger.info(
|
||||
f"Bulk sync completed: {result.created} created, "
|
||||
f"{result.updated} updated, {result.failed} failed"
|
||||
)
|
||||
return result
|
||||
except ValueError as e:
|
||||
logger.error(f"Bulk sync failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/groups",
|
||||
summary="List all groups",
|
||||
response_model=GroupsListResponse,
|
||||
responses={
|
||||
200: {"description": "List of groups"},
|
||||
},
|
||||
)
|
||||
async def list_groups(
|
||||
search: Optional[str] = Query(None, description="Search by group name"),
|
||||
offset: int = Query(0, ge=0, description="Number of records to skip"),
|
||||
limit: int = Query(50, ge=1, le=100, description="Maximum records to return"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> GroupsListResponse:
|
||||
"""
|
||||
List all groups synced from Authentik
|
||||
|
||||
Returns paginated list of groups with their details.
|
||||
Supports search filtering by name.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
items, total = await service.list_groups(
|
||||
search=search,
|
||||
offset=offset,
|
||||
limit=limit,
|
||||
)
|
||||
return GroupsListResponse(items=items, total=total)
|
||||
|
||||
@router.post(
|
||||
"/groups/sync-from-authentik",
|
||||
summary="Bulk sync groups from Authentik",
|
||||
response_model=BulkSyncResultSchema,
|
||||
responses={
|
||||
200: {"description": "Sync completed"},
|
||||
401: {"description": "Authentik API credentials invalid"},
|
||||
503: {"description": "Authentik service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_groups_from_authentik(
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all groups from Authentik and sync to local database
|
||||
|
||||
This endpoint uses the Authentik admin API to fetch all groups
|
||||
and create/update them in the local database. Requires
|
||||
AUTHENTIK_USERNAME and AUTHENTIK_PASSWORD to be configured.
|
||||
|
||||
Use this to populate groups or to re-sync after changes in Authentik.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
result = await service.bulk_sync_groups_from_authentik()
|
||||
logger.info(
|
||||
f"Groups bulk sync completed: {result.created} created, "
|
||||
f"{result.updated} updated, {result.failed} failed"
|
||||
)
|
||||
return result
|
||||
except ValueError as e:
|
||||
logger.error(f"Groups bulk sync failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/me",
|
||||
summary="Get current user profile",
|
||||
response_model=AuthSyncResponse,
|
||||
responses={
|
||||
200: {"description": "User profile"},
|
||||
401: {"description": "Not authenticated"},
|
||||
404: {"description": "User not found in database"},
|
||||
},
|
||||
)
|
||||
async def get_me(
|
||||
forward_auth_user: Optional[dict] = Depends(get_forward_auth_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> AuthSyncResponse:
|
||||
"""
|
||||
Get the current authenticated user's profile
|
||||
|
||||
Authentication is handled by NPM forward auth with Authentik.
|
||||
The proxy sets X-authentik-* headers which this endpoint reads.
|
||||
|
||||
For internal/LAN access (no forward auth headers), returns 401.
|
||||
Use POST /auth/sync with an OIDC token for mobile app authentication.
|
||||
"""
|
||||
# Require forward auth for this endpoint
|
||||
if forward_auth_user is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authentication required - access via authenticated proxy or use /auth/sync",
|
||||
)
|
||||
|
||||
service = AuthService(session)
|
||||
|
||||
# Try to find user by Authentik UID first, then by email
|
||||
user = None
|
||||
uid = forward_auth_user.get("uid")
|
||||
if uid:
|
||||
try:
|
||||
import uuid
|
||||
authentik_id = uuid.UUID(uid)
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
except (ValueError, TypeError):
|
||||
pass # Invalid UUID, try email
|
||||
|
||||
if user is None:
|
||||
email = forward_auth_user.get("email")
|
||||
if email:
|
||||
user = await service.get_user_by_email(email)
|
||||
|
||||
if user is None:
|
||||
# User authenticated with Authentik but not synced to database yet
|
||||
# This can happen on first login via web
|
||||
logger.info(f"User {forward_auth_user.get('email')} not found, creating from forward auth")
|
||||
|
||||
# Create user from forward auth headers
|
||||
from src.auth.schemas import TokenInfoSchema
|
||||
token_info = TokenInfoSchema(
|
||||
sub=forward_auth_user.get("uid", ""),
|
||||
email=forward_auth_user.get("email", ""),
|
||||
name=forward_auth_user.get("name"),
|
||||
groups=forward_auth_user.get("groups", []),
|
||||
)
|
||||
|
||||
try:
|
||||
user, _ = await service.sync_user(token_info)
|
||||
await service.sync_roles(user, token_info.groups)
|
||||
await session.commit()
|
||||
await session.refresh(user, ["preferences", "roles"])
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create user from forward auth: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="Failed to create user profile",
|
||||
)
|
||||
|
||||
# Sync roles from current groups (in case they changed)
|
||||
groups = forward_auth_user.get("groups", [])
|
||||
roles = await service.sync_roles(user, groups)
|
||||
await session.commit()
|
||||
|
||||
return AuthSyncResponse(
|
||||
user=service.user_to_schema(user),
|
||||
roles=service.roles_to_schema(roles),
|
||||
preferences=service.preferences_to_schema(user.preferences),
|
||||
is_new_user=False,
|
||||
)
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
auth_controller = AuthController()
|
||||
@@ -0,0 +1,129 @@
|
||||
"""
|
||||
Authentication Schemas
|
||||
|
||||
Pydantic models for auth request/response payloads.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from pydantic import Field
|
||||
|
||||
from src.base_schema import BaseSchema
|
||||
|
||||
|
||||
class AuthSyncRequest(BaseSchema):
|
||||
"""
|
||||
Request payload for POST /auth/sync
|
||||
|
||||
The client sends this after obtaining an OIDC token from Authentik.
|
||||
The access_token is validated against Authentik's userinfo endpoint.
|
||||
"""
|
||||
|
||||
access_token: str = Field(
|
||||
...,
|
||||
description="OIDC access token from Authentik",
|
||||
)
|
||||
|
||||
|
||||
class RoleSchema(BaseSchema):
|
||||
"""Role information in domain:action format"""
|
||||
|
||||
name: str = Field(..., description="Role name (e.g., 'control-room:admin')")
|
||||
domain: str = Field(..., description="Permission domain (e.g., 'control-room')")
|
||||
action: str = Field(..., description="Permission action (e.g., 'admin')")
|
||||
|
||||
|
||||
class UserPreferencesSchema(BaseSchema):
|
||||
"""User preferences"""
|
||||
|
||||
theme: str = Field(default="system", description="Theme preference: system, light, dark")
|
||||
default_room: str = Field(default="front-hall", description="Default room for housekeeping")
|
||||
preferences_json: dict = Field(default_factory=dict, description="Extended preferences")
|
||||
|
||||
|
||||
class UserSchema(BaseSchema):
|
||||
"""User information returned from sync"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal user ID")
|
||||
authentik_id: uuid.UUID = Field(..., description="Authentik user ID")
|
||||
email: str = Field(..., description="User email")
|
||||
name: str = Field(..., description="Display name")
|
||||
avatar_url: Optional[str] = Field(None, description="Profile picture URL")
|
||||
created_at: datetime = Field(..., description="Account creation timestamp")
|
||||
last_login: Optional[datetime] = Field(None, description="Last login timestamp")
|
||||
|
||||
|
||||
class AuthSyncResponse(BaseSchema):
|
||||
"""
|
||||
Response from POST /auth/sync
|
||||
|
||||
Contains the synced user profile, roles, and preferences.
|
||||
"""
|
||||
|
||||
user: UserSchema = Field(..., description="User profile")
|
||||
roles: list[RoleSchema] = Field(..., description="User's permission roles")
|
||||
preferences: UserPreferencesSchema = Field(..., description="User preferences")
|
||||
is_new_user: bool = Field(..., description="True if user was just created")
|
||||
|
||||
|
||||
class TokenInfoSchema(BaseSchema):
|
||||
"""
|
||||
Token information from Authentik userinfo endpoint
|
||||
|
||||
This is what Authentik returns when validating an access token.
|
||||
"""
|
||||
|
||||
sub: str = Field(..., description="Subject (Authentik user ID)")
|
||||
email: str = Field(..., description="User email")
|
||||
name: Optional[str] = Field(None, description="Display name")
|
||||
preferred_username: Optional[str] = Field(None, description="Username")
|
||||
groups: list[str] = Field(default_factory=list, description="Group memberships")
|
||||
picture: Optional[str] = Field(None, description="Profile picture URL")
|
||||
|
||||
|
||||
class UserListItemSchema(BaseSchema):
|
||||
"""User item for list display"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal user ID")
|
||||
email: str = Field(..., description="User email")
|
||||
name: str = Field(..., description="Display name")
|
||||
avatar_url: Optional[str] = Field(None, description="Profile picture URL")
|
||||
created_at: datetime = Field(..., description="Account creation timestamp")
|
||||
last_login: Optional[datetime] = Field(None, description="Last login timestamp")
|
||||
roles: list[str] = Field(default_factory=list, description="Role names")
|
||||
|
||||
|
||||
class UsersListResponse(BaseSchema):
|
||||
"""Response from GET /auth/users"""
|
||||
|
||||
items: list[UserListItemSchema] = Field(..., description="List of users")
|
||||
total: int = Field(..., description="Total count of users")
|
||||
|
||||
|
||||
class BulkSyncResultSchema(BaseSchema):
|
||||
"""Result from bulk sync operation"""
|
||||
|
||||
created: int = Field(..., description="Number of users created")
|
||||
updated: int = Field(..., description="Number of users updated")
|
||||
failed: int = Field(..., description="Number of users that failed to sync")
|
||||
total_in_authentik: int = Field(..., description="Total users in Authentik")
|
||||
errors: list[str] = Field(default_factory=list, description="Error messages for failed syncs")
|
||||
|
||||
|
||||
class GroupListItemSchema(BaseSchema):
|
||||
"""Group item for list display"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal group ID")
|
||||
authentik_id: uuid.UUID = Field(..., description="Authentik group ID")
|
||||
name: str = Field(..., description="Group name")
|
||||
is_superuser: bool = Field(default=False, description="Whether group has superuser privileges")
|
||||
parent_name: Optional[str] = Field(None, description="Parent group name")
|
||||
member_count: int = Field(default=0, description="Number of users in this group")
|
||||
synced_at: datetime = Field(..., description="Last sync timestamp")
|
||||
|
||||
|
||||
class GroupsListResponse(BaseSchema):
|
||||
"""Response from GET /auth/groups"""
|
||||
|
||||
items: list[GroupListItemSchema] = Field(..., description="List of groups")
|
||||
total: int = Field(..., description="Total count of groups")
|
||||
@@ -0,0 +1,667 @@
|
||||
"""
|
||||
Authentication Service
|
||||
|
||||
Business logic for user synchronization from Authentik.
|
||||
"""
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.config import get_settings
|
||||
from src.logging_config import get_logger
|
||||
from src.db.models import User, Role, UserPreferences, Group
|
||||
from src.auth.schemas import TokenInfoSchema, UserSchema, RoleSchema, UserPreferencesSchema, UserListItemSchema, BulkSyncResultSchema, GroupListItemSchema
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class AuthService:
|
||||
"""
|
||||
Service for authentication and user synchronization
|
||||
|
||||
Handles:
|
||||
- Token validation via Authentik userinfo endpoint
|
||||
- User creation/update from OIDC claims
|
||||
- Role synchronization from Authentik groups
|
||||
"""
|
||||
|
||||
def __init__(self, session: AsyncSession):
|
||||
"""
|
||||
Initialize auth service
|
||||
|
||||
Args:
|
||||
session: Async database session
|
||||
"""
|
||||
self.session = session
|
||||
self.userinfo_url = f"{settings.authentik_url}/application/o/userinfo/"
|
||||
|
||||
async def validate_token(self, access_token: str) -> TokenInfoSchema:
|
||||
"""
|
||||
Validate access token via Authentik userinfo endpoint
|
||||
|
||||
Args:
|
||||
access_token: OIDC access token
|
||||
|
||||
Returns:
|
||||
Token info containing user claims
|
||||
|
||||
Raises:
|
||||
ValueError: If token is invalid or expired
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
response = await client.get(
|
||||
self.userinfo_url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise ValueError("Invalid or expired token")
|
||||
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
logger.debug(f"Userinfo response: {data}")
|
||||
|
||||
return TokenInfoSchema(
|
||||
sub=data.get("sub"),
|
||||
email=data.get("email"),
|
||||
name=data.get("name") or data.get("preferred_username"),
|
||||
preferred_username=data.get("preferred_username"),
|
||||
groups=data.get("groups", []),
|
||||
picture=data.get("picture"),
|
||||
)
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
logger.error(f"Authentik userinfo request failed: {e}")
|
||||
raise ValueError(f"Token validation failed: {e.response.status_code}")
|
||||
except httpx.RequestError as e:
|
||||
logger.error(f"Authentik userinfo request error: {e}")
|
||||
raise ValueError("Authentication service unavailable")
|
||||
|
||||
async def get_user_by_email(self, email: str) -> Optional[User]:
|
||||
"""
|
||||
Get user by email address
|
||||
|
||||
Args:
|
||||
email: User email address
|
||||
|
||||
Returns:
|
||||
User if found, None otherwise
|
||||
"""
|
||||
stmt = (
|
||||
select(User)
|
||||
.options(selectinload(User.roles), selectinload(User.preferences))
|
||||
.where(User.email == email)
|
||||
)
|
||||
result = await self.session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_user_by_authentik_id(self, authentik_id: uuid.UUID) -> Optional[User]:
|
||||
"""
|
||||
Get user by Authentik UUID
|
||||
|
||||
Args:
|
||||
authentik_id: Authentik user UUID
|
||||
|
||||
Returns:
|
||||
User if found, None otherwise
|
||||
"""
|
||||
stmt = (
|
||||
select(User)
|
||||
.options(selectinload(User.roles), selectinload(User.preferences))
|
||||
.where(User.authentik_id == authentik_id)
|
||||
)
|
||||
result = await self.session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def sync_user(self, token_info: TokenInfoSchema) -> tuple[User, bool]:
|
||||
"""
|
||||
Create or update user from OIDC token info
|
||||
|
||||
Args:
|
||||
token_info: Validated token information
|
||||
|
||||
Returns:
|
||||
Tuple of (User, is_new_user)
|
||||
"""
|
||||
authentik_id = uuid.UUID(token_info.sub)
|
||||
|
||||
# Try to find existing user
|
||||
stmt = (
|
||||
select(User)
|
||||
.options(selectinload(User.roles), selectinload(User.preferences))
|
||||
.where(User.authentik_id == authentik_id)
|
||||
)
|
||||
result = await self.session.execute(stmt)
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
is_new = user is None
|
||||
|
||||
if is_new:
|
||||
# Create new user
|
||||
user = User(
|
||||
authentik_id=authentik_id,
|
||||
email=token_info.email,
|
||||
name=token_info.name or token_info.email,
|
||||
avatar_url=token_info.picture,
|
||||
last_login=datetime.now(timezone.utc),
|
||||
)
|
||||
self.session.add(user)
|
||||
await self.session.flush() # Get the user ID
|
||||
|
||||
# Create default preferences
|
||||
preferences = UserPreferences(user_id=user.id)
|
||||
self.session.add(preferences)
|
||||
|
||||
logger.info(f"Created new user: {token_info.email}")
|
||||
else:
|
||||
# Update existing user
|
||||
user.email = token_info.email
|
||||
user.name = token_info.name or token_info.email
|
||||
user.avatar_url = token_info.picture
|
||||
user.last_login = datetime.now(timezone.utc)
|
||||
|
||||
logger.info(f"Updated existing user: {token_info.email}")
|
||||
|
||||
await self.session.flush()
|
||||
return user, is_new
|
||||
|
||||
async def sync_roles(self, user: User, groups: list[str]) -> list[Role]:
|
||||
"""
|
||||
Synchronize user roles from Authentik groups
|
||||
|
||||
Maps Authentik groups (e.g., 'tatlock-control-room-admin')
|
||||
to application roles (e.g., 'control-room:admin').
|
||||
|
||||
Args:
|
||||
user: User to sync roles for
|
||||
groups: List of Authentik group names
|
||||
|
||||
Returns:
|
||||
List of synced Role objects
|
||||
"""
|
||||
# Get all roles that match the user's Authentik groups
|
||||
stmt = select(Role).where(Role.authentik_group.in_(groups))
|
||||
result = await self.session.execute(stmt)
|
||||
matching_roles = list(result.scalars().all())
|
||||
|
||||
# Clear existing roles and set new ones
|
||||
user.roles = matching_roles
|
||||
|
||||
role_names = [r.name for r in matching_roles]
|
||||
logger.info(f"Synced roles for {user.email}: {role_names}")
|
||||
|
||||
return matching_roles
|
||||
|
||||
def user_to_schema(self, user: User) -> UserSchema:
|
||||
"""Convert User model to schema"""
|
||||
return UserSchema(
|
||||
id=user.id,
|
||||
authentik_id=user.authentik_id,
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
avatar_url=user.avatar_url,
|
||||
created_at=user.created_at,
|
||||
last_login=user.last_login,
|
||||
)
|
||||
|
||||
def roles_to_schema(self, roles: list[Role]) -> list[RoleSchema]:
|
||||
"""Convert Role models to schemas"""
|
||||
return [
|
||||
RoleSchema(name=r.name, domain=r.domain, action=r.action)
|
||||
for r in roles
|
||||
]
|
||||
|
||||
def preferences_to_schema(self, preferences: Optional[UserPreferences]) -> UserPreferencesSchema:
|
||||
"""Convert UserPreferences model to schema"""
|
||||
if preferences is None:
|
||||
return UserPreferencesSchema()
|
||||
|
||||
return UserPreferencesSchema(
|
||||
theme=preferences.theme,
|
||||
default_room=preferences.default_room,
|
||||
preferences_json=preferences.preferences_json or {},
|
||||
)
|
||||
|
||||
async def list_users(
|
||||
self,
|
||||
search: Optional[str] = None,
|
||||
offset: int = 0,
|
||||
limit: int = 50,
|
||||
) -> tuple[list[UserListItemSchema], int]:
|
||||
"""
|
||||
List all users with optional search and pagination
|
||||
|
||||
Args:
|
||||
search: Optional search query (matches name or email)
|
||||
offset: Number of records to skip
|
||||
limit: Maximum number of records to return
|
||||
|
||||
Returns:
|
||||
Tuple of (list of user schemas, total count)
|
||||
"""
|
||||
from sqlalchemy import func
|
||||
|
||||
# Base query with roles loaded
|
||||
base_query = select(User).options(selectinload(User.roles))
|
||||
|
||||
# Apply search filter if provided
|
||||
if search:
|
||||
search_filter = f"%{search}%"
|
||||
base_query = base_query.where(
|
||||
(User.name.ilike(search_filter)) | (User.email.ilike(search_filter))
|
||||
)
|
||||
|
||||
# Get total count
|
||||
count_query = select(func.count()).select_from(base_query.subquery())
|
||||
total_result = await self.session.execute(count_query)
|
||||
total = total_result.scalar() or 0
|
||||
|
||||
# Apply pagination and ordering
|
||||
query = base_query.order_by(User.name).offset(offset).limit(limit)
|
||||
result = await self.session.execute(query)
|
||||
users = list(result.scalars().all())
|
||||
|
||||
# Convert to schemas
|
||||
items = [
|
||||
UserListItemSchema(
|
||||
id=user.id,
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
avatar_url=user.avatar_url,
|
||||
created_at=user.created_at,
|
||||
last_login=user.last_login,
|
||||
roles=[role.name for role in user.roles],
|
||||
)
|
||||
for user in users
|
||||
]
|
||||
|
||||
return items, total
|
||||
|
||||
def _extract_cookie(self, headers: httpx.Headers, cookie_name: str) -> str:
|
||||
"""Extract a specific cookie value from Set-Cookie headers"""
|
||||
for header in headers.get_list('set-cookie'):
|
||||
if header.startswith(f'{cookie_name}='):
|
||||
match = re.match(rf'{cookie_name}=([^;]+)', header)
|
||||
if match:
|
||||
return match.group(1)
|
||||
return ""
|
||||
|
||||
async def _authentik_session_login(self, client: httpx.AsyncClient) -> str:
|
||||
"""
|
||||
Authenticate with Authentik using the flow API to establish a session
|
||||
|
||||
Authentik's flow API requires:
|
||||
1. Cookie persistence between requests (manually handled due to domain restrictions)
|
||||
2. X-authentik-CSRF header set to the authentik_csrf cookie value
|
||||
3. Multi-stage flow handling (identification -> password -> done)
|
||||
|
||||
Args:
|
||||
client: httpx client
|
||||
|
||||
Returns:
|
||||
Session cookie value for subsequent API calls
|
||||
|
||||
Raises:
|
||||
ValueError: If authentication fails
|
||||
"""
|
||||
flow_url = f"{settings.authentik_url}/api/v3/flows/executor/default-authentication-flow/"
|
||||
|
||||
# Step 1: Get the initial flow challenge (this sets the session and csrf cookies)
|
||||
resp = await client.get(flow_url, headers={"Accept": "application/json"})
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Extract cookies manually from Set-Cookie headers (bypasses domain restrictions)
|
||||
session_cookie = self._extract_cookie(resp.headers, "authentik_session")
|
||||
csrf_cookie = self._extract_cookie(resp.headers, "authentik_csrf")
|
||||
|
||||
logger.debug(f"Flow initial: component={data.get('component')}, session={bool(session_cookie)}, csrf={bool(csrf_cookie)}")
|
||||
|
||||
# Build headers with manual cookie and CSRF token
|
||||
def build_headers():
|
||||
hdrs = {
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"Cookie": f"authentik_session={session_cookie}",
|
||||
}
|
||||
if csrf_cookie:
|
||||
hdrs["Cookie"] += f"; authentik_csrf={csrf_cookie}"
|
||||
hdrs["X-authentik-CSRF"] = csrf_cookie
|
||||
return hdrs
|
||||
|
||||
# Step 2: Handle identification stage - submit username
|
||||
if data.get("component") == "ak-stage-identification":
|
||||
resp = await client.post(
|
||||
flow_url,
|
||||
json={"uid_field": settings.authentik_username},
|
||||
headers=build_headers(),
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Update session cookie if new one received
|
||||
new_session = self._extract_cookie(resp.headers, "authentik_session")
|
||||
if new_session:
|
||||
session_cookie = new_session
|
||||
|
||||
logger.debug(f"After username: component={data.get('component')}")
|
||||
|
||||
# Step 3: Handle password stage if required
|
||||
if data.get("component") == "ak-stage-password":
|
||||
resp = await client.post(
|
||||
flow_url,
|
||||
json={"password": settings.authentik_password},
|
||||
headers=build_headers(),
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Update session cookie if new one received
|
||||
new_session = self._extract_cookie(resp.headers, "authentik_session")
|
||||
if new_session:
|
||||
session_cookie = new_session
|
||||
|
||||
logger.debug(f"After password: component={data.get('component')}")
|
||||
|
||||
# Check for access denied
|
||||
if data.get("component") == "ak-stage-access-denied":
|
||||
raise ValueError("Authentik authentication failed: access denied")
|
||||
|
||||
# Check for redirect (successful auth)
|
||||
if data.get("component") == "xak-flow-redirect" or data.get("to"):
|
||||
logger.info("Successfully authenticated with Authentik via flow")
|
||||
return session_cookie
|
||||
|
||||
# If we're still in identification stage, the username might be wrong
|
||||
if data.get("component") == "ak-stage-identification":
|
||||
response_errors = data.get("response_errors", {})
|
||||
raise ValueError(f"Authentication stuck at identification stage: {response_errors}")
|
||||
|
||||
logger.info(f"Authentik flow completed with component: {data.get('component')}")
|
||||
return session_cookie
|
||||
|
||||
async def bulk_sync_from_authentik(self) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all users from Authentik admin API and sync to local database
|
||||
|
||||
Returns:
|
||||
BulkSyncResultSchema with counts of created/updated/failed users
|
||||
"""
|
||||
if not settings.authentik_username or not settings.authentik_password:
|
||||
raise ValueError("AUTHENTIK_USERNAME and AUTHENTIK_PASSWORD must be configured")
|
||||
|
||||
created = 0
|
||||
updated = 0
|
||||
failed = 0
|
||||
errors = []
|
||||
total_in_authentik = 0
|
||||
|
||||
# Step 1: Fetch all user data from Authentik API
|
||||
authentik_users = []
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0, follow_redirects=True) as client:
|
||||
# Authenticate with Authentik to get session cookie
|
||||
session_cookie = await self._authentik_session_login(client)
|
||||
|
||||
# Fetch users from Authentik admin API using session cookie
|
||||
response = await client.get(
|
||||
f"{settings.authentik_url}/api/v3/core/users/",
|
||||
params={"page_size": 500},
|
||||
headers={
|
||||
"Accept": "application/json",
|
||||
"Cookie": f"authentik_session={session_cookie}",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise ValueError("Authentik API token is invalid or expired")
|
||||
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
authentik_users = data.get("results", [])
|
||||
total_in_authentik = data.get("pagination", {}).get("count", len(authentik_users))
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise ValueError(f"Authentik API error: {e.response.status_code}")
|
||||
except httpx.RequestError as e:
|
||||
raise ValueError(f"Failed to connect to Authentik: {str(e)}")
|
||||
|
||||
# Step 2: Sync users to database (outside of httpx context to avoid greenlet issues)
|
||||
for auth_user in authentik_users:
|
||||
try:
|
||||
# Skip service accounts and inactive users
|
||||
if auth_user.get("type") in ("service_account", "internal_service_account"):
|
||||
continue
|
||||
if not auth_user.get("is_active", True):
|
||||
continue
|
||||
|
||||
# Extract user data from Authentik
|
||||
authentik_id = uuid.UUID(auth_user["uuid"])
|
||||
email = auth_user.get("email") or f"{auth_user['username']}@local"
|
||||
name = auth_user.get("name") or auth_user.get("username", "Unknown")
|
||||
avatar_url = auth_user.get("avatar")
|
||||
|
||||
# Get user's groups for role mapping
|
||||
groups = []
|
||||
groups_summary = auth_user.get("groups_obj", [])
|
||||
for group in groups_summary:
|
||||
groups.append(group.get("name", ""))
|
||||
|
||||
# Check if user exists
|
||||
stmt = select(User).where(User.authentik_id == authentik_id)
|
||||
result = await self.session.execute(stmt)
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
if user is None:
|
||||
# Create new user
|
||||
user = User(
|
||||
authentik_id=authentik_id,
|
||||
email=email,
|
||||
name=name,
|
||||
avatar_url=avatar_url,
|
||||
)
|
||||
self.session.add(user)
|
||||
await self.session.flush()
|
||||
|
||||
# Create default preferences
|
||||
preferences = UserPreferences(user_id=user.id)
|
||||
self.session.add(preferences)
|
||||
created += 1
|
||||
logger.info(f"Created user from Authentik: {email}")
|
||||
else:
|
||||
# Update existing user
|
||||
user.email = email
|
||||
user.name = name
|
||||
user.avatar_url = avatar_url
|
||||
updated += 1
|
||||
logger.info(f"Updated user from Authentik: {email}")
|
||||
|
||||
# Sync roles from groups
|
||||
await self.sync_roles(user, groups)
|
||||
|
||||
except Exception as e:
|
||||
failed += 1
|
||||
error_msg = f"Failed to sync user {auth_user.get('username', 'unknown')}: {str(e)}"
|
||||
errors.append(error_msg)
|
||||
logger.warning(error_msg)
|
||||
|
||||
# Commit all changes
|
||||
await self.session.commit()
|
||||
|
||||
return BulkSyncResultSchema(
|
||||
created=created,
|
||||
updated=updated,
|
||||
failed=failed,
|
||||
total_in_authentik=total_in_authentik,
|
||||
errors=errors,
|
||||
)
|
||||
|
||||
async def list_groups(
|
||||
self,
|
||||
search: Optional[str] = None,
|
||||
offset: int = 0,
|
||||
limit: int = 50,
|
||||
) -> tuple[list[GroupListItemSchema], int]:
|
||||
"""
|
||||
List all groups with optional search and pagination
|
||||
|
||||
Args:
|
||||
search: Optional search query (matches name)
|
||||
offset: Number of records to skip
|
||||
limit: Maximum number of records to return
|
||||
|
||||
Returns:
|
||||
Tuple of (list of group schemas, total count)
|
||||
"""
|
||||
from sqlalchemy import func
|
||||
|
||||
# Base query
|
||||
base_query = select(Group)
|
||||
|
||||
# Apply search filter if provided
|
||||
if search:
|
||||
search_filter = f"%{search}%"
|
||||
base_query = base_query.where(Group.name.ilike(search_filter))
|
||||
|
||||
# Get total count
|
||||
count_query = select(func.count()).select_from(base_query.subquery())
|
||||
total_result = await self.session.execute(count_query)
|
||||
total = total_result.scalar() or 0
|
||||
|
||||
# Apply pagination and ordering
|
||||
query = base_query.order_by(Group.name).offset(offset).limit(limit)
|
||||
result = await self.session.execute(query)
|
||||
groups = list(result.scalars().all())
|
||||
|
||||
# Convert to schemas
|
||||
items = [
|
||||
GroupListItemSchema(
|
||||
id=group.id,
|
||||
authentik_id=group.authentik_id,
|
||||
name=group.name,
|
||||
is_superuser=group.is_superuser,
|
||||
parent_name=group.parent_name,
|
||||
member_count=group.member_count,
|
||||
synced_at=group.synced_at,
|
||||
)
|
||||
for group in groups
|
||||
]
|
||||
|
||||
return items, total
|
||||
|
||||
async def bulk_sync_groups_from_authentik(self) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all groups from Authentik admin API and sync to local database
|
||||
|
||||
Returns:
|
||||
BulkSyncResultSchema with counts of created/updated/failed groups
|
||||
"""
|
||||
if not settings.authentik_username or not settings.authentik_password:
|
||||
raise ValueError("AUTHENTIK_USERNAME and AUTHENTIK_PASSWORD must be configured")
|
||||
|
||||
created = 0
|
||||
updated = 0
|
||||
failed = 0
|
||||
errors = []
|
||||
total_in_authentik = 0
|
||||
|
||||
# Step 1: Fetch all group data from Authentik API
|
||||
authentik_groups = []
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0, follow_redirects=True) as client:
|
||||
# Authenticate with Authentik to get session cookie
|
||||
session_cookie = await self._authentik_session_login(client)
|
||||
|
||||
# Fetch groups from Authentik admin API using session cookie
|
||||
response = await client.get(
|
||||
f"{settings.authentik_url}/api/v3/core/groups/",
|
||||
params={"page_size": 500},
|
||||
headers={
|
||||
"Accept": "application/json",
|
||||
"Cookie": f"authentik_session={session_cookie}",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise ValueError("Authentik API token is invalid or expired")
|
||||
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
authentik_groups = data.get("results", [])
|
||||
total_in_authentik = data.get("pagination", {}).get("count", len(authentik_groups))
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise ValueError(f"Authentik API error: {e.response.status_code}")
|
||||
except httpx.RequestError as e:
|
||||
raise ValueError(f"Failed to connect to Authentik: {str(e)}")
|
||||
|
||||
# Step 2: Sync groups to database (outside of httpx context to avoid greenlet issues)
|
||||
for auth_group in authentik_groups:
|
||||
try:
|
||||
# Extract group data from Authentik
|
||||
authentik_id = uuid.UUID(auth_group["pk"])
|
||||
name = auth_group.get("name", "Unknown")
|
||||
is_superuser = auth_group.get("is_superuser", False)
|
||||
parent_name = auth_group.get("parent_name")
|
||||
# users field contains list of user PKs
|
||||
member_count = len(auth_group.get("users", []))
|
||||
|
||||
# Check if group exists
|
||||
stmt = select(Group).where(Group.authentik_id == authentik_id)
|
||||
result = await self.session.execute(stmt)
|
||||
group = result.scalar_one_or_none()
|
||||
|
||||
if group is None:
|
||||
# Create new group
|
||||
group = Group(
|
||||
authentik_id=authentik_id,
|
||||
name=name,
|
||||
is_superuser=is_superuser,
|
||||
parent_name=parent_name,
|
||||
member_count=member_count,
|
||||
)
|
||||
self.session.add(group)
|
||||
created += 1
|
||||
logger.info(f"Created group from Authentik: {name}")
|
||||
else:
|
||||
# Update existing group
|
||||
group.name = name
|
||||
group.is_superuser = is_superuser
|
||||
group.parent_name = parent_name
|
||||
group.member_count = member_count
|
||||
updated += 1
|
||||
logger.info(f"Updated group from Authentik: {name}")
|
||||
|
||||
except Exception as e:
|
||||
failed += 1
|
||||
error_msg = f"Failed to sync group {auth_group.get('name', 'unknown')}: {str(e)}"
|
||||
errors.append(error_msg)
|
||||
logger.warning(error_msg)
|
||||
|
||||
# Commit all changes
|
||||
await self.session.commit()
|
||||
|
||||
return BulkSyncResultSchema(
|
||||
created=created,
|
||||
updated=updated,
|
||||
failed=failed,
|
||||
total_in_authentik=total_in_authentik,
|
||||
errors=errors,
|
||||
)
|
||||
|
||||
|
||||
# Factory function for dependency injection
|
||||
def get_auth_service(session: AsyncSession) -> AuthService:
|
||||
"""Create AuthService instance with database session"""
|
||||
return AuthService(session)
|
||||
@@ -1,197 +0,0 @@
|
||||
"""
|
||||
Core-AI HTTP Client
|
||||
|
||||
Provides interface to Core-AI service for AI performance metrics.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Optional, Dict, List, Any
|
||||
from src.logging_config import get_logger
|
||||
from src.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class CoreAIClient:
|
||||
"""
|
||||
HTTP client for Core-AI service
|
||||
|
||||
Provides access to AI performance metrics, tool execution stats,
|
||||
and memory system monitoring.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
timeout: int = 10
|
||||
):
|
||||
"""
|
||||
Initialize Core-AI client
|
||||
|
||||
Args:
|
||||
base_url: Core-AI base URL (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or getattr(settings, 'core_ai_base_url', 'http://core-ai:8086')).rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.client = httpx.AsyncClient(timeout=self.timeout)
|
||||
|
||||
async def close(self):
|
||||
"""Close the HTTP client"""
|
||||
await self.client.aclose()
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if Core-AI service is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
response = await self.client.get(f"{self.base_url}/health")
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"Core-AI health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_metrics(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Get comprehensive AI performance metrics
|
||||
|
||||
Returns:
|
||||
Dict with agent performance, tool execution, memory stats
|
||||
|
||||
Example:
|
||||
{
|
||||
"uptime_seconds": 3600,
|
||||
"timestamp": "2025-12-03T20:00:00Z",
|
||||
"agent": {
|
||||
"total_requests": 100,
|
||||
"avg_response_time_ms": 1250.5,
|
||||
"p95_response_time_ms": 3200.0,
|
||||
...
|
||||
},
|
||||
"tools": {
|
||||
"total_calls": 250,
|
||||
"success_rate": 0.98,
|
||||
"top_tools": {...}
|
||||
},
|
||||
"memory": {
|
||||
"tier1_hit_rate": 0.85,
|
||||
...
|
||||
},
|
||||
...
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response = await self.client.get(f"{self.base_url}/metrics")
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except httpx.HTTPStatusError as e:
|
||||
logger.error(f"Failed to get metrics: HTTP {e.response.status_code}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get metrics: {e}")
|
||||
raise
|
||||
|
||||
async def get_recent_errors(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get recent request errors
|
||||
|
||||
Args:
|
||||
limit: Maximum number of errors to return
|
||||
|
||||
Returns:
|
||||
List of error records with timestamps
|
||||
|
||||
Example:
|
||||
[
|
||||
{
|
||||
"timestamp": "2025-12-03T19:45:12Z",
|
||||
"agent_type": "pydantic",
|
||||
"error": "Connection timeout",
|
||||
"duration_ms": 5000
|
||||
},
|
||||
...
|
||||
]
|
||||
"""
|
||||
try:
|
||||
response = await self.client.get(
|
||||
f"{self.base_url}/metrics/errors",
|
||||
params={"limit": limit}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data.get("errors", [])
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get recent errors: {e}")
|
||||
raise
|
||||
|
||||
async def get_tool_failures(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get recent tool execution failures
|
||||
|
||||
Args:
|
||||
limit: Maximum number of failures to return
|
||||
|
||||
Returns:
|
||||
List of tool failure records
|
||||
|
||||
Example:
|
||||
[
|
||||
{
|
||||
"timestamp": "2025-12-03T19:50:30Z",
|
||||
"tool_name": "list_containers",
|
||||
"error": "Connection refused",
|
||||
"duration_ms": 150
|
||||
},
|
||||
...
|
||||
]
|
||||
"""
|
||||
try:
|
||||
response = await self.client.get(
|
||||
f"{self.base_url}/metrics/tool-failures",
|
||||
params={"limit": limit}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data.get("failures", [])
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get tool failures: {e}")
|
||||
raise
|
||||
|
||||
async def reset_metrics(self) -> bool:
|
||||
"""
|
||||
Reset all metrics (admin operation)
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
try:
|
||||
response = await self.client.post(f"{self.base_url}/metrics/reset")
|
||||
response.raise_for_status()
|
||||
logger.info("Successfully reset Core-AI metrics")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to reset metrics: {e}")
|
||||
raise
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry"""
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Async context manager exit"""
|
||||
await self.close()
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_ai_client: Optional[CoreAIClient] = None
|
||||
|
||||
|
||||
def get_ai_client() -> CoreAIClient:
|
||||
"""Get singleton Core-AI client instance"""
|
||||
global _ai_client
|
||||
if _ai_client is None:
|
||||
_ai_client = CoreAIClient()
|
||||
return _ai_client
|
||||
@@ -0,0 +1,409 @@
|
||||
"""
|
||||
Home Assistant REST API Client
|
||||
|
||||
Provides interface to Home Assistant REST API for home automation control.
|
||||
Uses long-lived access token authentication.
|
||||
API Reference: https://developers.home-assistant.io/docs/api/rest/
|
||||
"""
|
||||
import httpx
|
||||
import json
|
||||
from typing import Optional, Dict, List, Any
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from src.logging_config import get_logger
|
||||
from src.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class HomeAssistantClient:
|
||||
"""
|
||||
HTTP client for Home Assistant REST API
|
||||
|
||||
Uses long-lived access token authentication via Bearer token.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
token: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize Home Assistant client
|
||||
|
||||
Args:
|
||||
base_url: Home Assistant base URL (default from settings)
|
||||
token: Long-lived access token (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.homeassistant_url).rstrip("/")
|
||||
self.token = token or settings.homeassistant_token
|
||||
self.timeout = timeout
|
||||
|
||||
if not self.token:
|
||||
logger.warning("Home Assistant token not configured")
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with Bearer token authentication"""
|
||||
return {
|
||||
"Authorization": f"Bearer {self.token}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
# ========================================================================
|
||||
# Health & Discovery
|
||||
# ========================================================================
|
||||
|
||||
async def health_check(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Check Home Assistant API connectivity and get version info
|
||||
|
||||
HA Endpoint: GET /api/
|
||||
|
||||
Returns:
|
||||
Dict with connected status, platform name, and version
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
return {
|
||||
"status": "healthy",
|
||||
"connected": True,
|
||||
"platform": "home_assistant",
|
||||
"version": data.get("version", "unknown")
|
||||
}
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"connected": False,
|
||||
"platform": "home_assistant",
|
||||
"error": f"HTTP {response.status_code}"
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Home Assistant health check failed: {e}")
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"connected": False,
|
||||
"platform": "home_assistant",
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
async def get_states(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get all entity states
|
||||
|
||||
HA Endpoint: GET /api/states
|
||||
|
||||
Returns:
|
||||
List of all entity states
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/states",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_state(self, entity_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Get state of a specific entity
|
||||
|
||||
HA Endpoint: GET /api/states/<entity_id>
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., "light.living_room")
|
||||
|
||||
Returns:
|
||||
Entity state dict or None if not found
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/states/{entity_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
if response.status_code == 404:
|
||||
return None
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_config(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Get Home Assistant configuration (includes areas)
|
||||
|
||||
HA Endpoint: GET /api/config
|
||||
|
||||
Returns:
|
||||
Configuration dict including components, location, etc.
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/config",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
# ========================================================================
|
||||
# Device Control
|
||||
# ========================================================================
|
||||
|
||||
async def call_service(
|
||||
self,
|
||||
domain: str,
|
||||
service: str,
|
||||
entity_id: Optional[str] = None,
|
||||
service_data: Optional[Dict[str, Any]] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Call a Home Assistant service
|
||||
|
||||
HA Endpoint: POST /api/services/<domain>/<service>
|
||||
|
||||
Args:
|
||||
domain: Service domain (e.g., "light", "switch", "scene")
|
||||
service: Service name (e.g., "turn_on", "turn_off", "toggle")
|
||||
entity_id: Target entity ID (optional for some services)
|
||||
service_data: Additional service data/attributes
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
payload = service_data.copy() if service_data else {}
|
||||
if entity_id:
|
||||
payload["entity_id"] = entity_id
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/services/{domain}/{service}",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def turn_on(
|
||||
self,
|
||||
entity_id: str,
|
||||
**attributes
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Turn on an entity with optional attributes
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., "light.living_room")
|
||||
**attributes: Additional attributes (brightness, color_temp, etc.)
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="turn_on",
|
||||
entity_id=entity_id,
|
||||
service_data=attributes if attributes else None
|
||||
)
|
||||
|
||||
async def turn_off(self, entity_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Turn off an entity
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="turn_off",
|
||||
entity_id=entity_id
|
||||
)
|
||||
|
||||
async def toggle(self, entity_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Toggle an entity
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="toggle",
|
||||
entity_id=entity_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Scenes
|
||||
# ========================================================================
|
||||
|
||||
async def activate_scene(self, scene_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Activate a scene
|
||||
|
||||
Args:
|
||||
scene_id: Scene entity ID (e.g., "scene.movie_night")
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="scene",
|
||||
service="turn_on",
|
||||
entity_id=scene_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Scripts
|
||||
# ========================================================================
|
||||
|
||||
async def run_script(
|
||||
self,
|
||||
script_id: str,
|
||||
variables: Optional[Dict[str, Any]] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Execute a script with optional variables
|
||||
|
||||
Args:
|
||||
script_id: Script entity ID (e.g., "script.bedtime_routine")
|
||||
variables: Script variables
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
service_data = {"variables": variables} if variables else None
|
||||
return await self.call_service(
|
||||
domain="script",
|
||||
service="turn_on",
|
||||
entity_id=script_id,
|
||||
service_data=service_data
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Automations
|
||||
# ========================================================================
|
||||
|
||||
async def enable_automation(self, automation_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Enable an automation
|
||||
|
||||
Args:
|
||||
automation_id: Automation entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="automation",
|
||||
service="turn_on",
|
||||
entity_id=automation_id
|
||||
)
|
||||
|
||||
async def disable_automation(self, automation_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Disable an automation
|
||||
|
||||
Args:
|
||||
automation_id: Automation entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="automation",
|
||||
service="turn_off",
|
||||
entity_id=automation_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# History
|
||||
# ========================================================================
|
||||
|
||||
async def get_history(
|
||||
self,
|
||||
entity_id: str,
|
||||
hours: int = 24
|
||||
) -> List[List[Dict[str, Any]]]:
|
||||
"""
|
||||
Get state history for an entity
|
||||
|
||||
HA Endpoint: GET /api/history/period/<timestamp>
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID to get history for
|
||||
hours: Number of hours of history (default 24)
|
||||
|
||||
Returns:
|
||||
List of state history entries
|
||||
"""
|
||||
start_time = datetime.now(timezone.utc) - timedelta(hours=hours)
|
||||
timestamp = start_time.isoformat()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/history/period/{timestamp}",
|
||||
headers=self._get_headers(),
|
||||
params={
|
||||
"filter_entity_id": entity_id,
|
||||
"minimal_response": "true"
|
||||
}
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
# ========================================================================
|
||||
# Areas (via template API)
|
||||
# ========================================================================
|
||||
|
||||
async def get_areas(self) -> List[Dict[str, str]]:
|
||||
"""
|
||||
Get all areas/rooms
|
||||
|
||||
Note: The REST API doesn't have a direct areas endpoint.
|
||||
This uses the template API to render area data.
|
||||
|
||||
HA Endpoint: POST /api/template
|
||||
|
||||
Returns:
|
||||
List of area dicts with id and name
|
||||
"""
|
||||
template = """
|
||||
{% set areas_list = [] %}
|
||||
{% for area in areas() %}
|
||||
{% set areas_list = areas_list + [{"id": area, "name": area_name(area)}] %}
|
||||
{% endfor %}
|
||||
{{ areas_list | tojson }}
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/template",
|
||||
headers=self._get_headers(),
|
||||
json={"template": template}
|
||||
)
|
||||
response.raise_for_status()
|
||||
# Response is rendered template as string
|
||||
return json.loads(response.text)
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_homeassistant_client: Optional[HomeAssistantClient] = None
|
||||
|
||||
|
||||
def get_homeassistant_client() -> HomeAssistantClient:
|
||||
"""Get singleton Home Assistant client instance"""
|
||||
global _homeassistant_client
|
||||
if _homeassistant_client is None:
|
||||
_homeassistant_client = HomeAssistantClient()
|
||||
return _homeassistant_client
|
||||
@@ -1,561 +0,0 @@
|
||||
"""
|
||||
Uptime Kuma Socket.IO Client
|
||||
|
||||
Provides interface to Uptime Kuma via Socket.IO for monitor management.
|
||||
Also provides metrics API access for real-time status data.
|
||||
"""
|
||||
import socketio
|
||||
import asyncio
|
||||
import httpx
|
||||
import re
|
||||
from typing import Optional, Dict, List, Any
|
||||
from src.logging_config import get_logger
|
||||
from src.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class KumaClient:
|
||||
"""
|
||||
Socket.IO client for Uptime Kuma
|
||||
|
||||
Uses Socket.IO for real-time communication with Uptime Kuma.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
password: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize Kuma client
|
||||
|
||||
Args:
|
||||
base_url: Kuma base URL (default from settings)
|
||||
username: Kuma username (default from settings)
|
||||
password: Kuma password (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.kuma_url).rstrip("/")
|
||||
self.username = username or settings.kuma_username
|
||||
self.password = password or settings.kuma_password
|
||||
self.timeout = timeout
|
||||
|
||||
self.sio = socketio.AsyncClient(
|
||||
reconnection=True,
|
||||
reconnection_attempts=3,
|
||||
reconnection_delay=1,
|
||||
)
|
||||
self._connected = False
|
||||
self._authenticated = False
|
||||
self._monitors_cache: Dict[int, Dict[str, Any]] = {}
|
||||
|
||||
if not self.username or not self.password:
|
||||
logger.warning("Uptime Kuma credentials not configured")
|
||||
|
||||
async def _ensure_connected(self):
|
||||
"""Ensure we have an active connection and authentication"""
|
||||
if not self._connected:
|
||||
await self.connect()
|
||||
if not self._authenticated:
|
||||
await self.login()
|
||||
|
||||
async def connect(self):
|
||||
"""Connect to Uptime Kuma Socket.IO server"""
|
||||
if self._connected:
|
||||
return
|
||||
|
||||
try:
|
||||
await self.sio.connect(self.base_url, transports=['websocket'])
|
||||
self._connected = True
|
||||
logger.info(f"Connected to Uptime Kuma at {self.base_url}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to Uptime Kuma: {e}")
|
||||
raise
|
||||
|
||||
async def disconnect(self):
|
||||
"""Disconnect from Uptime Kuma"""
|
||||
if self._connected:
|
||||
await self.sio.disconnect()
|
||||
self._connected = False
|
||||
self._authenticated = False
|
||||
logger.info("Disconnected from Uptime Kuma")
|
||||
|
||||
async def login(self):
|
||||
"""Authenticate with Uptime Kuma"""
|
||||
if not self._connected:
|
||||
await self.connect()
|
||||
|
||||
try:
|
||||
# Uptime Kuma login event
|
||||
login_response = await self.sio.call(
|
||||
'login',
|
||||
{
|
||||
'username': self.username,
|
||||
'password': self.password,
|
||||
'token': None
|
||||
},
|
||||
timeout=self.timeout
|
||||
)
|
||||
|
||||
if login_response and login_response.get('ok'):
|
||||
self._authenticated = True
|
||||
logger.info("Successfully authenticated with Uptime Kuma")
|
||||
else:
|
||||
error_msg = login_response.get('msg', 'Unknown error') if login_response else 'No response'
|
||||
raise Exception(f"Login failed: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to authenticate with Uptime Kuma: {e}")
|
||||
raise
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if Uptime Kuma is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
await self._ensure_connected()
|
||||
return self._authenticated
|
||||
except Exception as e:
|
||||
logger.error(f"Uptime Kuma health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_monitors(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all monitors with uptime data
|
||||
|
||||
Returns:
|
||||
List of monitor configurations with uptime_24h field
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Storage for monitor list and uptime data received via events
|
||||
monitor_list_data = {}
|
||||
uptime_list_data = {}
|
||||
monitor_event_received = asyncio.Event()
|
||||
uptime_event_received = asyncio.Event()
|
||||
|
||||
# Register event handler for monitorList
|
||||
@self.sio.event
|
||||
async def monitorList(data):
|
||||
nonlocal monitor_list_data
|
||||
monitor_list_data = data
|
||||
monitor_event_received.set()
|
||||
|
||||
# Register event handler for uptimeList (24h uptime percentages)
|
||||
@self.sio.event
|
||||
async def uptimeList(monitor_id, uptime_data):
|
||||
nonlocal uptime_list_data
|
||||
# uptime_data is typically a dict with time periods: {"24": 99.5, "720": 98.2, ...}
|
||||
uptime_list_data[str(monitor_id)] = uptime_data
|
||||
# Don't set event here as we'll get multiple calls
|
||||
|
||||
# Request monitor list - this triggers the server to send monitorList event
|
||||
response = await self.sio.call('getMonitorList', timeout=self.timeout)
|
||||
logger.info(f"getMonitorList call response: {response}")
|
||||
|
||||
# Wait for the monitorList event (with timeout)
|
||||
try:
|
||||
await asyncio.wait_for(monitor_event_received.wait(), timeout=5.0)
|
||||
logger.info(f"Received monitorList event with {len(monitor_list_data)} items")
|
||||
|
||||
# Give time for uptimeList events to arrive
|
||||
await asyncio.sleep(0.5)
|
||||
logger.info(f"Received uptime data for {len(uptime_list_data)} monitors")
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning("Timeout waiting for monitorList event")
|
||||
|
||||
# Process the monitor list data
|
||||
if monitor_list_data and isinstance(monitor_list_data, dict):
|
||||
monitors = []
|
||||
for monitor_id, monitor_data in monitor_list_data.items():
|
||||
if isinstance(monitor_data, dict):
|
||||
monitor_data['id'] = int(monitor_id)
|
||||
|
||||
# Add uptime data if available
|
||||
uptime_info = uptime_list_data.get(str(monitor_id), {})
|
||||
if isinstance(uptime_info, dict):
|
||||
# Uptime Kuma provides 24h uptime as key "24"
|
||||
monitor_data['uptime_24h'] = float(uptime_info.get('24', 0))
|
||||
else:
|
||||
monitor_data['uptime_24h'] = 0.0
|
||||
|
||||
monitors.append(monitor_data)
|
||||
self._monitors_cache[int(monitor_id)] = monitor_data
|
||||
|
||||
logger.info(f"Found {len(monitors)} monitors total")
|
||||
return monitors
|
||||
|
||||
logger.warning(f"No valid monitor data received")
|
||||
return []
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get monitors: {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
async def get_monitor(self, monitor_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
Get details of a specific monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
Monitor configuration details
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
response = await self.sio.call('getMonitor', monitor_id, timeout=self.timeout)
|
||||
|
||||
if response:
|
||||
self._monitors_cache[monitor_id] = response
|
||||
return response
|
||||
|
||||
raise Exception(f"Monitor {monitor_id} not found")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get monitor {monitor_id}: {e}")
|
||||
raise
|
||||
|
||||
async def find_monitor_by_name(self, name: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Find a monitor by its name (case-insensitive)
|
||||
|
||||
Args:
|
||||
name: Monitor name to search for
|
||||
|
||||
Returns:
|
||||
Monitor object if found, None otherwise
|
||||
"""
|
||||
monitors = await self.get_monitors()
|
||||
name_lower = name.lower()
|
||||
|
||||
for monitor in monitors:
|
||||
if monitor.get("name", "").lower() == name_lower:
|
||||
return monitor
|
||||
|
||||
return None
|
||||
|
||||
async def find_monitors_by_tag(self, tag: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Find all monitors with a specific tag
|
||||
|
||||
Args:
|
||||
tag: Tag name to search for
|
||||
|
||||
Returns:
|
||||
List of monitors with the tag
|
||||
"""
|
||||
monitors = await self.get_monitors()
|
||||
tagged_monitors = []
|
||||
|
||||
for monitor in monitors:
|
||||
monitor_tags = monitor.get("tags", [])
|
||||
if any(t.get("name", "").lower() == tag.lower() for t in monitor_tags):
|
||||
tagged_monitors.append(monitor)
|
||||
|
||||
return tagged_monitors
|
||||
|
||||
async def pause_monitor(self, monitor_id: int) -> bool:
|
||||
"""
|
||||
Pause a monitor (disable monitoring)
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Uptime Kuma pause event
|
||||
response = await self.sio.call('pauseMonitor', monitor_id, timeout=self.timeout)
|
||||
|
||||
if response and response.get('ok'):
|
||||
logger.info(f"Paused monitor {monitor_id}")
|
||||
return True
|
||||
|
||||
error_msg = response.get('msg', 'Unknown error') if response else 'No response'
|
||||
raise Exception(f"Failed to pause monitor: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to pause monitor {monitor_id}: {e}")
|
||||
raise
|
||||
|
||||
async def resume_monitor(self, monitor_id: int) -> bool:
|
||||
"""
|
||||
Resume a monitor (enable monitoring)
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Uptime Kuma resume event
|
||||
response = await self.sio.call('resumeMonitor', monitor_id, timeout=self.timeout)
|
||||
|
||||
if response and response.get('ok'):
|
||||
logger.info(f"Resumed monitor {monitor_id}")
|
||||
return True
|
||||
|
||||
error_msg = response.get('msg', 'Unknown error') if response else 'No response'
|
||||
raise Exception(f"Failed to resume monitor: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to resume monitor {monitor_id}: {e}")
|
||||
raise
|
||||
|
||||
async def pause_monitor_by_name(self, name: str) -> bool:
|
||||
"""
|
||||
Pause a monitor by its name
|
||||
|
||||
Args:
|
||||
name: Monitor name
|
||||
|
||||
Returns:
|
||||
True if successful, False if monitor not found
|
||||
"""
|
||||
monitor = await self.find_monitor_by_name(name)
|
||||
if not monitor:
|
||||
logger.warning(f"Monitor '{name}' not found")
|
||||
return False
|
||||
|
||||
await self.pause_monitor(monitor["id"])
|
||||
return True
|
||||
|
||||
async def resume_monitor_by_name(self, name: str) -> bool:
|
||||
"""
|
||||
Resume a monitor by its name
|
||||
|
||||
Args:
|
||||
name: Monitor name
|
||||
|
||||
Returns:
|
||||
True if successful, False if monitor not found
|
||||
"""
|
||||
monitor = await self.find_monitor_by_name(name)
|
||||
if not monitor:
|
||||
logger.warning(f"Monitor '{name}' not found")
|
||||
return False
|
||||
|
||||
await self.resume_monitor(monitor["id"])
|
||||
return True
|
||||
|
||||
async def add_monitor(self, monitor_config: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a new monitor
|
||||
|
||||
Args:
|
||||
monitor_config: Monitor configuration dict
|
||||
|
||||
Returns:
|
||||
Created monitor details including ID
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Uptime Kuma add monitor event
|
||||
response = await self.sio.call('add', monitor_config, timeout=self.timeout)
|
||||
|
||||
if response and response.get('ok'):
|
||||
monitor_id = response.get('monitorID')
|
||||
logger.info(f"Created monitor '{monitor_config.get('name')}' with ID {monitor_id}")
|
||||
|
||||
# Get full monitor details
|
||||
monitor = await self.get_monitor(monitor_id)
|
||||
return monitor
|
||||
|
||||
error_msg = response.get('msg', 'Unknown error') if response else 'No response'
|
||||
raise Exception(f"Failed to create monitor: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create monitor '{monitor_config.get('name')}': {e}")
|
||||
raise
|
||||
|
||||
async def update_monitor(self, monitor_id: int, monitor_config: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Update an existing monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
monitor_config: Updated monitor configuration
|
||||
|
||||
Returns:
|
||||
Updated monitor details
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Ensure ID is in the config
|
||||
monitor_config['id'] = monitor_id
|
||||
|
||||
# Uptime Kuma edit monitor event
|
||||
response = await self.sio.call('editMonitor', monitor_config, timeout=self.timeout)
|
||||
|
||||
if response and response.get('ok'):
|
||||
logger.info(f"Updated monitor {monitor_id}")
|
||||
|
||||
# Get updated monitor details
|
||||
monitor = await self.get_monitor(monitor_id)
|
||||
return monitor
|
||||
|
||||
error_msg = response.get('msg', 'Unknown error') if response else 'No response'
|
||||
raise Exception(f"Failed to update monitor: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update monitor {monitor_id}: {e}")
|
||||
raise
|
||||
|
||||
async def delete_monitor(self, monitor_id: int) -> bool:
|
||||
"""
|
||||
Delete a monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
await self._ensure_connected()
|
||||
|
||||
try:
|
||||
# Uptime Kuma delete monitor event
|
||||
response = await self.sio.call('deleteMonitor', monitor_id, timeout=self.timeout)
|
||||
|
||||
if response and response.get('ok'):
|
||||
logger.info(f"Deleted monitor {monitor_id}")
|
||||
|
||||
# Remove from cache
|
||||
self._monitors_cache.pop(monitor_id, None)
|
||||
return True
|
||||
|
||||
error_msg = response.get('msg', 'Unknown error') if response else 'No response'
|
||||
raise Exception(f"Failed to delete monitor: {error_msg}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to delete monitor {monitor_id}: {e}")
|
||||
raise
|
||||
|
||||
async def delete_monitor_by_name(self, name: str) -> bool:
|
||||
"""
|
||||
Delete a monitor by its name
|
||||
|
||||
Args:
|
||||
name: Monitor name
|
||||
|
||||
Returns:
|
||||
True if successful, False if monitor not found
|
||||
"""
|
||||
monitor = await self.find_monitor_by_name(name)
|
||||
if not monitor:
|
||||
logger.warning(f"Monitor '{name}' not found")
|
||||
return False
|
||||
|
||||
await self.delete_monitor(monitor["id"])
|
||||
return True
|
||||
|
||||
async def get_metrics_status(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""
|
||||
Get monitor status from Prometheus metrics endpoint
|
||||
|
||||
This is simpler and more reliable than Socket.IO for getting current status.
|
||||
Returns real-time UP/DOWN status but not historical uptime percentages.
|
||||
|
||||
Returns:
|
||||
Dict mapping monitor names to status info:
|
||||
{
|
||||
"Portainer": {
|
||||
"status": 1, # 1=UP, 0=DOWN, 2=PENDING, 3=MAINTENANCE
|
||||
"response_time": 5, # ms
|
||||
"monitor_type": "http",
|
||||
"url": "http://192.168.86.149:8001"
|
||||
},
|
||||
...
|
||||
}
|
||||
"""
|
||||
try:
|
||||
# Use API key authentication
|
||||
api_key = settings.kuma_api_key
|
||||
if not api_key:
|
||||
logger.warning("Kuma API key not configured")
|
||||
return {}
|
||||
|
||||
# Fetch metrics with HTTP Basic Auth (empty username, API key as password)
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/metrics",
|
||||
auth=("", api_key)
|
||||
)
|
||||
response.raise_for_status()
|
||||
metrics_text = response.text
|
||||
|
||||
# Parse Prometheus format metrics
|
||||
# Format: metric_name{label1="value1",label2="value2"} value
|
||||
monitor_data = {}
|
||||
|
||||
# Parse monitor_status lines
|
||||
status_pattern = r'monitor_status\{monitor_name="([^"]+)",.*?\} (\d+)'
|
||||
for match in re.finditer(status_pattern, metrics_text):
|
||||
monitor_name = match.group(1)
|
||||
status = int(match.group(2))
|
||||
|
||||
if monitor_name not in monitor_data:
|
||||
monitor_data[monitor_name] = {}
|
||||
monitor_data[monitor_name]['status'] = status
|
||||
|
||||
# Parse monitor_response_time lines
|
||||
response_pattern = r'monitor_response_time\{monitor_name="([^"]+)",monitor_type="([^"]+)",monitor_url="([^"]+)",.*?\} ([\d.]+)'
|
||||
for match in re.finditer(response_pattern, metrics_text):
|
||||
monitor_name = match.group(1)
|
||||
monitor_type = match.group(2)
|
||||
monitor_url = match.group(3)
|
||||
response_time = float(match.group(4))
|
||||
|
||||
if monitor_name not in monitor_data:
|
||||
monitor_data[monitor_name] = {}
|
||||
monitor_data[monitor_name].update({
|
||||
'response_time': response_time,
|
||||
'monitor_type': monitor_type,
|
||||
'url': monitor_url
|
||||
})
|
||||
|
||||
logger.info(f"Fetched metrics for {len(monitor_data)} monitors")
|
||||
return monitor_data
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch metrics: {e}")
|
||||
return {}
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry"""
|
||||
await self._ensure_connected()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Async context manager exit"""
|
||||
await self.disconnect()
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_kuma_client: Optional[KumaClient] = None
|
||||
|
||||
|
||||
def get_kuma_client() -> KumaClient:
|
||||
"""Get singleton Kuma client instance"""
|
||||
global _kuma_client
|
||||
if _kuma_client is None:
|
||||
_kuma_client = KumaClient()
|
||||
return _kuma_client
|
||||
+168
-112
@@ -210,6 +210,152 @@ class PortainerClient:
|
||||
response.raise_for_status()
|
||||
return True
|
||||
|
||||
async def get_stack_file(self, stack_id: int) -> str:
|
||||
"""
|
||||
Get the compose file content for a stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
|
||||
Returns:
|
||||
Docker Compose YAML content as string
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks/{stack_id}/file",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data.get("StackFileContent", "")
|
||||
|
||||
async def redeploy_stack(
|
||||
self,
|
||||
stack_id: int,
|
||||
endpoint_id: int,
|
||||
pull_image: bool = False
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Redeploy a stack with its current configuration
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
pull_image: Pull latest images before deployment
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
# Get current stack file content
|
||||
stack_content = await self.get_stack_file(stack_id)
|
||||
|
||||
# Get current stack to preserve env vars
|
||||
stack = await self.get_stack(stack_id)
|
||||
env_vars = stack.get("Env", [])
|
||||
|
||||
payload = {
|
||||
"stackFileContent": stack_content,
|
||||
"env": env_vars,
|
||||
"prune": False,
|
||||
"pullImage": pull_image
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def update_stack_env(
|
||||
self,
|
||||
stack_id: int,
|
||||
endpoint_id: int,
|
||||
env_vars: List[Dict[str, str]]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update stack environment variables
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
env_vars: List of {"name": "VAR_NAME", "value": "var_value"} dicts
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
# Get current stack file content (required for update)
|
||||
stack_content = await self.get_stack_file(stack_id)
|
||||
|
||||
payload = {
|
||||
"stackFileContent": stack_content,
|
||||
"env": env_vars,
|
||||
"prune": False,
|
||||
"pullImage": False
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def delete_container(
|
||||
self,
|
||||
endpoint_id: int,
|
||||
container_id: str,
|
||||
force: bool = False
|
||||
) -> bool:
|
||||
"""
|
||||
Delete a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
force: Force remove running container
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
params = {"force": "true" if force else "false"}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.delete(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}",
|
||||
headers=self._get_headers(),
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Deleted container {container_id}")
|
||||
return True
|
||||
|
||||
async def restart_container(self, endpoint_id: int, container_id: str) -> bool:
|
||||
"""
|
||||
Restart a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}/restart",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Restarted container {container_id}")
|
||||
return True
|
||||
|
||||
async def get_containers(self, endpoint_id: int, all_containers: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List containers on a specific endpoint
|
||||
@@ -292,70 +438,14 @@ class PortainerClient:
|
||||
return True
|
||||
|
||||
# ========================================================================
|
||||
# Docker Socket Fallback (for containers not managed by Portainer)
|
||||
# ========================================================================
|
||||
|
||||
async def _list_containers_via_socket(self, all_containers: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Fallback: List containers directly via Docker socket
|
||||
|
||||
Used when Portainer API doesn't return complete data (e.g., containers
|
||||
started outside Portainer, AMP game servers, etc.)
|
||||
|
||||
Args:
|
||||
all_containers: Include stopped containers
|
||||
|
||||
Returns:
|
||||
List of container details in Docker API format
|
||||
"""
|
||||
try:
|
||||
# Docker socket is mounted at /var/run/docker.sock
|
||||
# Use httpx with unix socket transport
|
||||
transport = httpx.AsyncHTTPTransport(uds="/var/run/docker.sock")
|
||||
async with httpx.AsyncClient(transport=transport, timeout=10) as client:
|
||||
params = {"all": 1 if all_containers else 0}
|
||||
response = await client.get(
|
||||
"http://localhost/v1.41/containers/json",
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as e:
|
||||
logger.warning(f"Docker socket fallback failed: {e}")
|
||||
return []
|
||||
|
||||
async def _inspect_container_via_socket(self, container_id_or_name: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Fallback: Inspect container directly via Docker socket
|
||||
|
||||
Args:
|
||||
container_id_or_name: Container ID or name
|
||||
|
||||
Returns:
|
||||
Container details or None
|
||||
"""
|
||||
try:
|
||||
transport = httpx.AsyncHTTPTransport(uds="/var/run/docker.sock")
|
||||
async with httpx.AsyncClient(transport=transport, timeout=10) as client:
|
||||
response = await client.get(
|
||||
f"http://localhost/v1.41/containers/{container_id_or_name}/json"
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as e:
|
||||
logger.warning(f"Docker socket inspect fallback failed for '{container_id_or_name}': {e}")
|
||||
return None
|
||||
|
||||
# ========================================================================
|
||||
# Helper methods for agent tools (auto-detect endpoint + fallback)
|
||||
# Helper methods for agent tools (auto-detect endpoint)
|
||||
# ========================================================================
|
||||
|
||||
async def list_containers(self, all_containers: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List containers using auto-detected endpoint with Docker socket fallback
|
||||
List containers using auto-detected endpoint
|
||||
|
||||
This is a convenience wrapper that automatically uses the first/default endpoint.
|
||||
If Portainer doesn't have complete data, falls back to Docker socket.
|
||||
|
||||
Args:
|
||||
all_containers: Include stopped containers (default: True)
|
||||
@@ -363,34 +453,18 @@ class PortainerClient:
|
||||
Returns:
|
||||
List of container details
|
||||
"""
|
||||
try:
|
||||
# Try Portainer first
|
||||
endpoints = await self.get_endpoints()
|
||||
if endpoints:
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
containers = await self.get_containers(endpoint_id, all_containers)
|
||||
if containers:
|
||||
return containers
|
||||
endpoints = await self.get_endpoints()
|
||||
if not endpoints:
|
||||
raise RuntimeError("No Portainer endpoints available")
|
||||
|
||||
# Fallback to Docker socket
|
||||
logger.info("Portainer returned no containers, trying Docker socket fallback...")
|
||||
return await self._list_containers_via_socket(all_containers)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing containers: {e}")
|
||||
# Try fallback even on exception
|
||||
try:
|
||||
return await self._list_containers_via_socket(all_containers)
|
||||
except Exception as fallback_error:
|
||||
logger.error(f"Fallback also failed: {fallback_error}")
|
||||
return []
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
return await self.get_containers(endpoint_id, all_containers)
|
||||
|
||||
async def inspect_container(self, container_name: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Inspect a container by name using auto-detected endpoint with Docker socket fallback
|
||||
Inspect a container by name using auto-detected endpoint
|
||||
|
||||
This is a convenience wrapper that automatically uses the first/default endpoint.
|
||||
If Portainer doesn't find the container, falls back to Docker socket.
|
||||
|
||||
Args:
|
||||
container_name: Container name (e.g., "jellyfin", "ollama")
|
||||
@@ -398,44 +472,26 @@ class PortainerClient:
|
||||
Returns:
|
||||
Container details or None if not found
|
||||
"""
|
||||
try:
|
||||
# Try Portainer first
|
||||
endpoints = await self.get_endpoints()
|
||||
if endpoints:
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
endpoints = await self.get_endpoints()
|
||||
if not endpoints:
|
||||
raise RuntimeError("No Portainer endpoints available")
|
||||
|
||||
# First list all containers to find the one matching the name
|
||||
all_containers = await self.get_containers(endpoint_id, all_containers=True)
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
|
||||
matching_container = None
|
||||
for container in all_containers:
|
||||
# Container names come as array like ['/jellyfin']
|
||||
names = container.get('Names', [])
|
||||
for name in names:
|
||||
clean_name = name.lstrip('/')
|
||||
if clean_name == container_name or clean_name.lower() == container_name.lower():
|
||||
matching_container = container
|
||||
break
|
||||
if matching_container:
|
||||
break
|
||||
# List all containers to find the one matching the name
|
||||
all_containers = await self.get_containers(endpoint_id, all_containers=True)
|
||||
|
||||
if matching_container:
|
||||
for container in all_containers:
|
||||
# Container names come as array like ['/jellyfin']
|
||||
names = container.get('Names', [])
|
||||
for name in names:
|
||||
clean_name = name.lstrip('/')
|
||||
if clean_name == container_name or clean_name.lower() == container_name.lower():
|
||||
# Get detailed info using container ID
|
||||
container_id = matching_container['Id']
|
||||
container_id = container['Id']
|
||||
return await self.get_container(endpoint_id, container_id)
|
||||
|
||||
# Not found in Portainer, try Docker socket fallback
|
||||
logger.info(f"Container '{container_name}' not found in Portainer, trying Docker socket fallback...")
|
||||
return await self._inspect_container_via_socket(container_name)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error inspecting container '{container_name}': {e}")
|
||||
# Try fallback even on exception
|
||||
try:
|
||||
return await self._inspect_container_via_socket(container_name)
|
||||
except Exception as fallback_error:
|
||||
logger.error(f"Fallback also failed: {fallback_error}")
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
# Singleton instance
|
||||
|
||||
+36
-49
@@ -1,5 +1,8 @@
|
||||
"""
|
||||
Global configuration for Core Code API
|
||||
|
||||
All configuration is loaded from environment variables or .env file.
|
||||
See .env.example for available settings.
|
||||
"""
|
||||
import tomllib
|
||||
from pathlib import Path
|
||||
@@ -20,31 +23,6 @@ def _get_version_from_pyproject() -> str:
|
||||
|
||||
__version__ = _get_version_from_pyproject()
|
||||
|
||||
# Import infrastructure credentials from gitignored module
|
||||
try:
|
||||
from src.credentials import (
|
||||
PORTAINER_URL, PORTAINER_API_KEY,
|
||||
NPM_URL, NPM_EMAIL, NPM_PASSWORD,
|
||||
KUMA_URL, KUMA_USERNAME, KUMA_PASSWORD, KUMA_API_KEY,
|
||||
BRAVE_SEARCH_API_KEY,
|
||||
GOOGLE_SEARCH_API_KEY, GOOGLE_SEARCH_ENGINE_ID
|
||||
)
|
||||
except ImportError:
|
||||
# Fallback to empty strings if credentials.py doesn't exist
|
||||
# (e.g., fresh clone before credentials setup)
|
||||
PORTAINER_URL = "http://localhost:8001"
|
||||
PORTAINER_API_KEY = ""
|
||||
NPM_URL = "http://localhost:81"
|
||||
NPM_EMAIL = ""
|
||||
NPM_PASSWORD = ""
|
||||
KUMA_URL = "http://localhost:3001"
|
||||
KUMA_USERNAME = ""
|
||||
KUMA_PASSWORD = ""
|
||||
KUMA_API_KEY = ""
|
||||
BRAVE_SEARCH_API_KEY = ""
|
||||
GOOGLE_SEARCH_API_KEY = ""
|
||||
GOOGLE_SEARCH_ENGINE_ID = ""
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Global application settings"""
|
||||
@@ -68,15 +46,13 @@ class Settings(BaseSettings):
|
||||
log_level: str = "DEBUG"
|
||||
|
||||
# Ollama Configuration (for AI orchestration)
|
||||
ollama_base_url: str = "http://ollama:11434"
|
||||
ollama_base_url: str # Required - set OLLAMA_BASE_URL in .env
|
||||
ollama_timeout: int = 300 # 5 minutes
|
||||
|
||||
# Model Configuration
|
||||
default_model: str = "mistral-tools:7b"
|
||||
agent_model: str = "gemma2:9b-instruct-q5_K_M" # Must support tool calling with ADK (~4GB VRAM)
|
||||
lightweight_models: str = "gemma3-tools:1b,phi3:mini"
|
||||
heavy_models: str = "mistral:7b,gemma2:9b,gemma3:12b,mixtral:8x7b"
|
||||
code_models: str = "codestral:latest,codegemma:latest"
|
||||
default_model: str = "mistral-nemo-large:latest"
|
||||
agent_model: str = "mistral-nemo-large:latest" # Must support tool calling with ADK (~4GB VRAM)
|
||||
code_models: str = "mistral-nemo-large:latest"
|
||||
# Previous config (gemma3:12b used ~10GB VRAM)
|
||||
# default_model: str = "gemma3:12b"
|
||||
# agent_model: str = "gemma3:12b"
|
||||
@@ -111,35 +87,45 @@ class Settings(BaseSettings):
|
||||
embedding_batch_size: int = 32
|
||||
|
||||
# Search Configuration
|
||||
search_provider: str = "google" # Options: google, brave, searxng, duckduckgo
|
||||
searxng_url: str = "http://searxng:8080" # For future self-hosted SearxNG
|
||||
search_provider: str = "searxng"
|
||||
searxng_url: str # Required - set SEARXNG_URL in .env
|
||||
|
||||
# Search API Keys (from credentials.py)
|
||||
brave_search_api_key: str = BRAVE_SEARCH_API_KEY # https://brave.com/search/api/
|
||||
google_search_api_key: str = GOOGLE_SEARCH_API_KEY # https://console.cloud.google.com/
|
||||
google_search_engine_id: str = GOOGLE_SEARCH_ENGINE_ID # Custom Search Engine ID
|
||||
# Infrastructure Management (Portainer)
|
||||
portainer_url: str # Required - set PORTAINER_URL in .env
|
||||
portainer_api_key: str # Required - set PORTAINER_API_KEY in .env
|
||||
|
||||
# Infrastructure Management (from credentials.py)
|
||||
portainer_url: str = PORTAINER_URL
|
||||
portainer_api_key: str = PORTAINER_API_KEY
|
||||
# Infrastructure Management (Nginx Proxy Manager)
|
||||
npm_url: str # Required - set NPM_URL in .env
|
||||
npm_email: str # Required - set NPM_EMAIL in .env
|
||||
npm_password: str # Required - set NPM_PASSWORD in .env
|
||||
|
||||
npm_url: str = NPM_URL
|
||||
npm_email: str = NPM_EMAIL
|
||||
npm_password: str = NPM_PASSWORD
|
||||
# Home Assistant Configuration
|
||||
homeassistant_url: str # Required - set HOMEASSISTANT_URL in .env
|
||||
homeassistant_token: str # Required - set HOMEASSISTANT_TOKEN in .env
|
||||
homeassistant_timeout: int = 30
|
||||
|
||||
kuma_url: str = KUMA_URL
|
||||
kuma_username: str = KUMA_USERNAME
|
||||
kuma_password: str = KUMA_PASSWORD
|
||||
kuma_api_key: str = KUMA_API_KEY
|
||||
# PostgreSQL Database
|
||||
postgres_host: str # Required - set POSTGRES_HOST in .env (e.g., localhost:5432)
|
||||
postgres_user: str = "core_api"
|
||||
postgres_password: str # Required - set POSTGRES_PASSWORD in .env
|
||||
postgres_database: str = "core_api"
|
||||
|
||||
# Core-AI Service (AI performance metrics)
|
||||
core_ai_base_url: str = "http://core-ai:8086"
|
||||
@property
|
||||
def database_url(self) -> str:
|
||||
"""Construct database URL from components"""
|
||||
return f"postgresql://{self.postgres_user}:{self.postgres_password}@{self.postgres_host}/{self.postgres_database}"
|
||||
|
||||
# OIDC Authentication (Authentik)
|
||||
oidc_enabled: bool = False # Set to True to require authentication
|
||||
oidc_issuer: str = "https://auth.schweitz.net/application/o/core-api/"
|
||||
oidc_audience: str = "core-api"
|
||||
|
||||
# Authentik API (for token validation and user management)
|
||||
# Must use domain name (not IP) when AUTHENTIK_COOKIE_DOMAIN is set
|
||||
authentik_url: str = "https://auth.schweitz.net" # Authentik base URL
|
||||
authentik_username: str = "" # Admin username for API access (AUTHENTIK_USERNAME env var)
|
||||
authentik_password: str = "" # Admin password for API access (AUTHENTIK_PASSWORD env var)
|
||||
|
||||
@property
|
||||
def model_aliases(self) -> dict:
|
||||
"""Computed property for model aliases"""
|
||||
@@ -165,6 +151,7 @@ class Settings(BaseSettings):
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = False
|
||||
extra = "ignore" # Ignore extra env vars not defined in Settings
|
||||
|
||||
|
||||
@lru_cache()
|
||||
|
||||
@@ -1,219 +0,0 @@
|
||||
"""
|
||||
AI Metrics Proxy Controller
|
||||
|
||||
Provides proxy endpoints to Core-AI service metrics.
|
||||
Allows external access to AI performance stats via core-api.
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from typing import Dict, List, Any
|
||||
from src.clients.ai_client import get_ai_client
|
||||
from src.logging_config import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Create router
|
||||
router = APIRouter(
|
||||
prefix="/ai",
|
||||
tags=["AI Metrics"]
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
summary="Check Core-AI service health",
|
||||
description="Verify that the Core-AI service is accessible and responding"
|
||||
)
|
||||
async def ai_health_check():
|
||||
"""
|
||||
Check if Core-AI service is healthy
|
||||
|
||||
Returns:
|
||||
Health status and availability
|
||||
"""
|
||||
try:
|
||||
ai_client = get_ai_client()
|
||||
is_healthy = await ai_client.health_check()
|
||||
|
||||
return {
|
||||
"service": "core-ai",
|
||||
"status": "healthy" if is_healthy else "unhealthy",
|
||||
"accessible": is_healthy
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"AI health check failed: {e}")
|
||||
return {
|
||||
"service": "core-ai",
|
||||
"status": "error",
|
||||
"accessible": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/metrics",
|
||||
response_model=Dict[str, Any],
|
||||
summary="Get comprehensive AI performance metrics",
|
||||
description="Returns detailed metrics including agent performance, tool execution stats, memory system metrics, and user activity"
|
||||
)
|
||||
async def get_ai_metrics():
|
||||
"""
|
||||
Proxy endpoint for Core-AI metrics
|
||||
|
||||
Returns comprehensive AI performance data:
|
||||
- Agent request statistics (total, by type, response times)
|
||||
- Response time percentiles (p50, p95, p99)
|
||||
- Tool execution metrics (calls, success rates, durations)
|
||||
- Memory system statistics (cache hits, consolidations)
|
||||
- User activity tracking
|
||||
- Concurrency metrics
|
||||
|
||||
Returns:
|
||||
Dict with all collected metrics
|
||||
|
||||
Raises:
|
||||
HTTPException: If Core-AI is unreachable or returns error
|
||||
"""
|
||||
try:
|
||||
ai_client = get_ai_client()
|
||||
metrics = await ai_client.get_metrics()
|
||||
return metrics
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch AI metrics: {e}")
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=f"Core-AI service unavailable: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/metrics/errors",
|
||||
response_model=Dict[str, Any],
|
||||
summary="Get recent request errors",
|
||||
description="Returns recent AI agent request errors with timestamps and details"
|
||||
)
|
||||
async def get_ai_errors(limit: int = 20):
|
||||
"""
|
||||
Get recent AI request errors
|
||||
|
||||
Args:
|
||||
limit: Maximum number of errors to return (default: 20)
|
||||
|
||||
Returns:
|
||||
Dict with error list and total count
|
||||
|
||||
Example response:
|
||||
{
|
||||
"errors": [
|
||||
{
|
||||
"timestamp": "2025-12-03T19:45:12Z",
|
||||
"agent_type": "pydantic",
|
||||
"error": "Connection timeout",
|
||||
"duration_ms": 5000
|
||||
}
|
||||
],
|
||||
"total": 1
|
||||
}
|
||||
"""
|
||||
try:
|
||||
ai_client = get_ai_client()
|
||||
errors = await ai_client.get_recent_errors(limit=limit)
|
||||
|
||||
return {
|
||||
"errors": errors,
|
||||
"total": len(errors)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch AI errors: {e}")
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=f"Core-AI service unavailable: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/metrics/tool-failures",
|
||||
response_model=Dict[str, Any],
|
||||
summary="Get recent tool execution failures",
|
||||
description="Returns recent tool execution failures with error details"
|
||||
)
|
||||
async def get_ai_tool_failures(limit: int = 20):
|
||||
"""
|
||||
Get recent tool execution failures
|
||||
|
||||
Args:
|
||||
limit: Maximum number of failures to return (default: 20)
|
||||
|
||||
Returns:
|
||||
Dict with failure list and total count
|
||||
|
||||
Example response:
|
||||
{
|
||||
"failures": [
|
||||
{
|
||||
"timestamp": "2025-12-03T19:50:30Z",
|
||||
"tool_name": "list_containers",
|
||||
"error": "Connection refused",
|
||||
"duration_ms": 150
|
||||
}
|
||||
],
|
||||
"total": 1
|
||||
}
|
||||
"""
|
||||
try:
|
||||
ai_client = get_ai_client()
|
||||
failures = await ai_client.get_tool_failures(limit=limit)
|
||||
|
||||
return {
|
||||
"failures": failures,
|
||||
"total": len(failures)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch tool failures: {e}")
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=f"Core-AI service unavailable: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/metrics/reset",
|
||||
summary="Reset all AI metrics (admin)",
|
||||
description="Clear all collected metrics. This is an administrative operation that resets all counters and history."
|
||||
)
|
||||
async def reset_ai_metrics():
|
||||
"""
|
||||
Reset all AI metrics (admin operation)
|
||||
|
||||
Clears all collected metrics including:
|
||||
- Request history
|
||||
- Tool execution stats
|
||||
- Memory system metrics
|
||||
- Error logs
|
||||
|
||||
Returns:
|
||||
Success confirmation
|
||||
|
||||
Note:
|
||||
This is an administrative operation that should be used carefully.
|
||||
All historical data will be lost.
|
||||
"""
|
||||
try:
|
||||
ai_client = get_ai_client()
|
||||
await ai_client.reset_metrics()
|
||||
|
||||
logger.info("AI metrics reset successfully")
|
||||
return {
|
||||
"success": True,
|
||||
"message": "AI metrics reset successfully"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to reset AI metrics: {e}")
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=f"Core-AI service unavailable: {str(e)}"
|
||||
)
|
||||
@@ -10,10 +10,8 @@ from src.controllers.base import BaseController
|
||||
from src.config import get_settings
|
||||
from src.logging_config import get_logger
|
||||
from src.models.ollama_client import get_ollama_client
|
||||
from src.db import get_database
|
||||
|
||||
# Note: Agent functionality moved to separate core-ai service (Dec 2025)
|
||||
# This service (core-api) only provides infrastructure management and tools
|
||||
AGENT_AVAILABLE = False
|
||||
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -52,20 +50,7 @@ class HealthController(BaseController):
|
||||
"service": settings.app_name,
|
||||
"version": settings.app_version,
|
||||
"status": "healthy",
|
||||
"documentation": {
|
||||
"swagger_ui": "/docs",
|
||||
"redoc": "/redoc",
|
||||
"openapi_spec": "/openapi.json"
|
||||
},
|
||||
"endpoints": {
|
||||
"chat_completions": "/v1/chat/completions",
|
||||
"models": "/v1/models",
|
||||
"conversations": "/v1/conversations",
|
||||
"web_scraper": "/web-scraper/scrape",
|
||||
"infrastructure": "/infrastructure",
|
||||
"health": "/health",
|
||||
"health_full": "/health/full"
|
||||
}
|
||||
"docs": "/docs"
|
||||
}
|
||||
|
||||
@router.get(
|
||||
@@ -75,17 +60,15 @@ class HealthController(BaseController):
|
||||
)
|
||||
async def health_check():
|
||||
"""
|
||||
Simple health check endpoint for container orchestration
|
||||
Fast health check endpoint for container orchestration
|
||||
|
||||
Returns a 200 OK status when the service is running properly.
|
||||
Used by Docker, Kubernetes, and load balancers.
|
||||
Returns a 200 OK immediately if the service is running.
|
||||
Does NOT check backend connectivity (use /health/full for that).
|
||||
Used by Docker, Kubernetes, and load balancers for liveness probes.
|
||||
"""
|
||||
ollama_client = get_ollama_client()
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
|
||||
return {
|
||||
"status": "healthy",
|
||||
"ollama_connected": ollama_healthy
|
||||
"version": settings.app_version
|
||||
}
|
||||
|
||||
@router.get(
|
||||
@@ -146,12 +129,18 @@ class HealthController(BaseController):
|
||||
ollama_error = str(e)
|
||||
logger.warning(f"Ollama health check failed: {ollama_error}")
|
||||
|
||||
# Note: Agent functionality moved to separate core-ai service
|
||||
# This service only needs Ollama for embeddings (infrastructure tools)
|
||||
# Agent health is checked separately in core-ai service
|
||||
# Check 2: Database connection
|
||||
database = get_database()
|
||||
db_healthy = False
|
||||
db_error = None
|
||||
|
||||
# Determine overall status (only Ollama required for core-api)
|
||||
is_healthy = ollama_healthy
|
||||
try:
|
||||
db_healthy = await database.health_check()
|
||||
except Exception as e:
|
||||
db_error = str(e)
|
||||
logger.warning(f"Database health check failed: {db_error}")
|
||||
|
||||
is_healthy = ollama_healthy and db_healthy
|
||||
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
status_code = 200 if is_healthy else 503
|
||||
@@ -170,7 +159,10 @@ class HealthController(BaseController):
|
||||
},
|
||||
"error": ollama_error
|
||||
},
|
||||
"note": "AI agent functionality available in separate core-ai service (port 8086)"
|
||||
"database": {
|
||||
"status": "✅ healthy" if db_healthy else "❌ unhealthy",
|
||||
"error": db_error
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -216,20 +208,7 @@ class HealthController(BaseController):
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
# 2. Agent Stack - Moved to separate core-ai service
|
||||
diagnostics["components"]["agent"] = {
|
||||
"status": "N/A",
|
||||
"note": "AI agent functionality moved to separate core-ai service (port 8086)",
|
||||
"check_url": "http://core-ai:8086/health"
|
||||
}
|
||||
|
||||
# 3. Memory System (Qdrant) - Moved to core-ai service
|
||||
diagnostics["components"]["qdrant"] = {
|
||||
"status": "N/A",
|
||||
"note": "Memory system managed by core-ai service (port 8086)"
|
||||
}
|
||||
|
||||
# 4. Configuration
|
||||
# 2. Configuration
|
||||
diagnostics["configuration"] = {
|
||||
"agent_fallback_enabled": settings.agent_fallback_enabled,
|
||||
"memory_tier1_max_turns": settings.memory_tier1_max_turns,
|
||||
|
||||
@@ -0,0 +1,748 @@
|
||||
"""
|
||||
Housekeeping Controller
|
||||
|
||||
Provides API endpoints for home automation via Home Assistant.
|
||||
Designed for the Tatlock Housekeeper agent and other consumers.
|
||||
"""
|
||||
import asyncio
|
||||
from fastapi import APIRouter, HTTPException, Query, Depends
|
||||
from typing import List, Dict, Any, Optional
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from src.controllers.base import BaseController
|
||||
from src.clients.homeassistant_client import get_homeassistant_client
|
||||
from src.logging_config import get_logger
|
||||
from src.auth.oidc import get_admin_user
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Pydantic Schemas
|
||||
# ========================================================================
|
||||
|
||||
class Device(BaseModel):
|
||||
"""Device/entity information"""
|
||||
entity_id: str
|
||||
name: str
|
||||
domain: str
|
||||
area: Optional[str] = None
|
||||
state: str
|
||||
attributes: Dict[str, Any] = {}
|
||||
last_changed: Optional[str] = None
|
||||
|
||||
|
||||
class DeviceListResponse(BaseModel):
|
||||
"""Response for device listing"""
|
||||
devices: List[Device]
|
||||
|
||||
|
||||
class DeviceDetailResponse(Device):
|
||||
"""Detailed device response"""
|
||||
pass
|
||||
|
||||
|
||||
class Area(BaseModel):
|
||||
"""Area/room information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class AreaListResponse(BaseModel):
|
||||
"""Response for area listing"""
|
||||
areas: List[Area]
|
||||
|
||||
|
||||
class DeviceControlRequest(BaseModel):
|
||||
"""Request to control a device"""
|
||||
action: str = Field(..., description="Action: turn_on, turn_off, or toggle")
|
||||
brightness: Optional[int] = Field(None, ge=0, le=255)
|
||||
color_temp: Optional[int] = None
|
||||
rgb_color: Optional[List[int]] = None
|
||||
|
||||
class Config:
|
||||
extra = "allow" # Allow additional attributes
|
||||
|
||||
|
||||
class DeviceControlResponse(BaseModel):
|
||||
"""Response from device control"""
|
||||
success: bool
|
||||
entity_id: str
|
||||
new_state: Optional[str] = None
|
||||
message: str
|
||||
|
||||
|
||||
class Scene(BaseModel):
|
||||
"""Scene information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class SceneListResponse(BaseModel):
|
||||
"""Response for scene listing"""
|
||||
scenes: List[Scene]
|
||||
|
||||
|
||||
class SceneActivateResponse(BaseModel):
|
||||
"""Response from scene activation"""
|
||||
success: bool
|
||||
scene_id: str
|
||||
message: str
|
||||
|
||||
|
||||
class Script(BaseModel):
|
||||
"""Script information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class ScriptListResponse(BaseModel):
|
||||
"""Response for script listing"""
|
||||
scripts: List[Script]
|
||||
|
||||
|
||||
class ScriptRunRequest(BaseModel):
|
||||
"""Request to run a script"""
|
||||
variables: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ScriptRunResponse(BaseModel):
|
||||
"""Response from script execution"""
|
||||
success: bool
|
||||
script_id: str
|
||||
message: str
|
||||
|
||||
|
||||
class Automation(BaseModel):
|
||||
"""Automation information"""
|
||||
id: str
|
||||
name: str
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutomationListResponse(BaseModel):
|
||||
"""Response for automation listing"""
|
||||
automations: List[Automation]
|
||||
|
||||
|
||||
class AutomationToggleRequest(BaseModel):
|
||||
"""Request to toggle automation"""
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutomationToggleResponse(BaseModel):
|
||||
"""Response from automation toggle"""
|
||||
success: bool
|
||||
automation_id: str
|
||||
enabled: bool
|
||||
message: str
|
||||
|
||||
|
||||
class HistoryEntry(BaseModel):
|
||||
"""Single history entry"""
|
||||
state: str
|
||||
timestamp: str
|
||||
attributes: Dict[str, Any] = {}
|
||||
|
||||
|
||||
class HistoryResponse(BaseModel):
|
||||
"""Response for history query"""
|
||||
entity_id: str
|
||||
history: List[HistoryEntry]
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""Health check response"""
|
||||
status: str
|
||||
connected: bool
|
||||
platform: str
|
||||
version: Optional[str] = None
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
class ErrorResponse(BaseModel):
|
||||
"""Standard error response"""
|
||||
error: bool = True
|
||||
code: str
|
||||
message: str
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Controller
|
||||
# ========================================================================
|
||||
|
||||
class HousekeepingController(BaseController):
|
||||
"""
|
||||
Controller for home automation operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Device discovery and control
|
||||
- Scene activation
|
||||
- Script execution
|
||||
- Automation management
|
||||
- State history
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/housekeeping", tags=["Housekeeping"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
# ====================================================================
|
||||
# Health
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
response_model=HealthResponse,
|
||||
summary="Home automation health check"
|
||||
)
|
||||
async def get_health():
|
||||
"""
|
||||
Check Home Assistant connection health
|
||||
|
||||
Returns connection status and HA version.
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
return await ha.health_check()
|
||||
|
||||
# ====================================================================
|
||||
# Device Discovery
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/devices",
|
||||
response_model=DeviceListResponse,
|
||||
summary="List available devices"
|
||||
)
|
||||
async def list_devices(
|
||||
domain: Optional[str] = Query(None, description="Filter by domain (light, switch, climate, etc.)"),
|
||||
area: Optional[str] = Query(None, description="Filter by area/room name")
|
||||
):
|
||||
"""
|
||||
List all available devices with optional filtering
|
||||
|
||||
Query Parameters:
|
||||
- domain: Filter by device type (light, switch, climate, media_player, etc.)
|
||||
- area: Filter by area/room name
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
# Non-controllable domains to filter out
|
||||
excluded_domains = {
|
||||
"zone", "person", "device_tracker", "sun", "weather",
|
||||
"persistent_notification", "update", "binary_sensor", "sensor",
|
||||
"conversation", "calendar", "button", "number", "select",
|
||||
"text", "time", "date", "datetime", "image", "tts", "stt"
|
||||
}
|
||||
|
||||
devices = []
|
||||
for state in states:
|
||||
entity_id = state.get("entity_id", "")
|
||||
entity_domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||
|
||||
# Skip non-controllable entities
|
||||
if entity_domain in excluded_domains:
|
||||
continue
|
||||
|
||||
# Apply domain filter
|
||||
if domain and entity_domain != domain:
|
||||
continue
|
||||
|
||||
# Get area from attributes
|
||||
device_area = state.get("attributes", {}).get("area_id")
|
||||
|
||||
# Apply area filter
|
||||
if area and device_area and area.lower() not in device_area.lower():
|
||||
continue
|
||||
|
||||
device = Device(
|
||||
entity_id=entity_id,
|
||||
name=state.get("attributes", {}).get("friendly_name", entity_id),
|
||||
domain=entity_domain,
|
||||
area=device_area,
|
||||
state=state.get("state", "unknown"),
|
||||
attributes=state.get("attributes", {}),
|
||||
last_changed=state.get("last_changed")
|
||||
)
|
||||
devices.append(device)
|
||||
|
||||
return DeviceListResponse(devices=devices)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list devices: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/devices/{entity_id:path}",
|
||||
response_model=DeviceDetailResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def get_device(entity_id: str):
|
||||
"""
|
||||
Get detailed state of a specific device
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., light.living_room)
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
state = await ha.get_state(entity_id)
|
||||
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "DEVICE_NOT_FOUND",
|
||||
"message": f"Device {entity_id} not found"}
|
||||
)
|
||||
|
||||
entity_domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||
|
||||
return DeviceDetailResponse(
|
||||
entity_id=entity_id,
|
||||
name=state.get("attributes", {}).get("friendly_name", entity_id),
|
||||
domain=entity_domain,
|
||||
area=state.get("attributes", {}).get("area_id"),
|
||||
state=state.get("state", "unknown"),
|
||||
attributes=state.get("attributes", {}),
|
||||
last_changed=state.get("last_changed")
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get device {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/areas",
|
||||
response_model=AreaListResponse,
|
||||
summary="List areas/rooms"
|
||||
)
|
||||
async def list_areas():
|
||||
"""
|
||||
List all configured areas/rooms in Home Assistant
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
areas = await ha.get_areas()
|
||||
return AreaListResponse(
|
||||
areas=[Area(id=a["id"], name=a["name"]) for a in areas]
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list areas: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
# ====================================================================
|
||||
# Device Control
|
||||
# ====================================================================
|
||||
|
||||
@router.post(
|
||||
"/devices/{entity_id:path}/control",
|
||||
response_model=DeviceControlResponse,
|
||||
responses={404: {"model": ErrorResponse}, 400: {"model": ErrorResponse}}
|
||||
)
|
||||
async def control_device(
|
||||
entity_id: str,
|
||||
request: DeviceControlRequest,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Control a device (turn on, turn off, toggle, or set attributes)
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., light.living_room)
|
||||
request: Control request with action and optional attributes
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
# Validate action
|
||||
valid_actions = ["turn_on", "turn_off", "toggle"]
|
||||
if request.action not in valid_actions:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": True, "code": "INVALID_ACTION",
|
||||
"message": f"Invalid action '{request.action}'. Must be one of: {', '.join(valid_actions)}"}
|
||||
)
|
||||
|
||||
try:
|
||||
# Check device exists first
|
||||
current_state = await ha.get_state(entity_id)
|
||||
if not current_state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "DEVICE_NOT_FOUND",
|
||||
"message": f"Device {entity_id} not found"}
|
||||
)
|
||||
|
||||
# Build attributes dict from request
|
||||
attributes = {}
|
||||
if request.brightness is not None:
|
||||
attributes["brightness"] = request.brightness
|
||||
if request.color_temp is not None:
|
||||
attributes["color_temp"] = request.color_temp
|
||||
if request.rgb_color is not None:
|
||||
attributes["rgb_color"] = request.rgb_color
|
||||
|
||||
# Add any extra attributes from request
|
||||
extra_fields = request.model_dump(exclude={"action", "brightness", "color_temp", "rgb_color"})
|
||||
for key, value in extra_fields.items():
|
||||
if value is not None:
|
||||
attributes[key] = value
|
||||
|
||||
# Execute action
|
||||
if request.action == "turn_on":
|
||||
await ha.turn_on(entity_id, **attributes)
|
||||
elif request.action == "turn_off":
|
||||
await ha.turn_off(entity_id)
|
||||
else: # toggle
|
||||
await ha.toggle(entity_id)
|
||||
|
||||
# Wait for HA to update state before fetching
|
||||
await asyncio.sleep(0.3)
|
||||
|
||||
# Get new state
|
||||
new_state = await ha.get_state(entity_id)
|
||||
|
||||
logger.info(f"Device {entity_id} controlled: {request.action} by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return DeviceControlResponse(
|
||||
success=True,
|
||||
entity_id=entity_id,
|
||||
new_state=new_state.get("state") if new_state else None,
|
||||
message=f"Device {request.action} successful"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to control device {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
# ====================================================================
|
||||
# Scenes
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/scenes",
|
||||
response_model=SceneListResponse,
|
||||
summary="List available scenes"
|
||||
)
|
||||
async def list_scenes():
|
||||
"""
|
||||
List all available scenes in Home Assistant
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
scenes = [
|
||||
Scene(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"])
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("scene.")
|
||||
]
|
||||
|
||||
return SceneListResponse(scenes=scenes)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list scenes: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/scenes/{scene_id:path}/activate",
|
||||
response_model=SceneActivateResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def activate_scene(
|
||||
scene_id: str,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Activate a scene
|
||||
|
||||
Args:
|
||||
scene_id: Scene entity ID (e.g., scene.movie_night)
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
# Verify scene exists
|
||||
state = await ha.get_state(scene_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "SCENE_NOT_FOUND",
|
||||
"message": f"Scene {scene_id} not found"}
|
||||
)
|
||||
|
||||
await ha.activate_scene(scene_id)
|
||||
|
||||
logger.info(f"Scene {scene_id} activated by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return SceneActivateResponse(
|
||||
success=True,
|
||||
scene_id=scene_id,
|
||||
message="Scene activated"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to activate scene {scene_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
# ====================================================================
|
||||
# Scripts
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/scripts",
|
||||
response_model=ScriptListResponse,
|
||||
summary="List available scripts"
|
||||
)
|
||||
async def list_scripts():
|
||||
"""
|
||||
List all available scripts/sequences in Home Assistant
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
scripts = [
|
||||
Script(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"])
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("script.")
|
||||
]
|
||||
|
||||
return ScriptListResponse(scripts=scripts)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list scripts: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/scripts/{script_id:path}/run",
|
||||
response_model=ScriptRunResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def run_script(
|
||||
script_id: str,
|
||||
request: Optional[ScriptRunRequest] = None,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Execute a script with optional variables
|
||||
|
||||
Args:
|
||||
script_id: Script entity ID (e.g., script.bedtime_routine)
|
||||
request: Optional variables for the script
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
# Verify script exists
|
||||
state = await ha.get_state(script_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "SCRIPT_NOT_FOUND",
|
||||
"message": f"Script {script_id} not found"}
|
||||
)
|
||||
|
||||
variables = request.variables if request else None
|
||||
await ha.run_script(script_id, variables)
|
||||
|
||||
logger.info(f"Script {script_id} executed by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return ScriptRunResponse(
|
||||
success=True,
|
||||
script_id=script_id,
|
||||
message="Script executed"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to run script {script_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
# ====================================================================
|
||||
# Automations
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/automations",
|
||||
response_model=AutomationListResponse,
|
||||
summary="List automations"
|
||||
)
|
||||
async def list_automations():
|
||||
"""
|
||||
List all automations with their enabled/disabled status
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
automations = [
|
||||
Automation(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"]),
|
||||
enabled=s.get("state") == "on"
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("automation.")
|
||||
]
|
||||
|
||||
return AutomationListResponse(automations=automations)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list automations: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/automations/{automation_id:path}/toggle",
|
||||
response_model=AutomationToggleResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def toggle_automation(
|
||||
automation_id: str,
|
||||
request: AutomationToggleRequest,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Enable or disable an automation
|
||||
|
||||
Args:
|
||||
automation_id: Automation entity ID (e.g., automation.motion_lights)
|
||||
request: Contains enabled boolean
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
# Verify automation exists
|
||||
state = await ha.get_state(automation_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "AUTOMATION_NOT_FOUND",
|
||||
"message": f"Automation {automation_id} not found"}
|
||||
)
|
||||
|
||||
if request.enabled:
|
||||
await ha.enable_automation(automation_id)
|
||||
else:
|
||||
await ha.disable_automation(automation_id)
|
||||
|
||||
logger.info(f"Automation {automation_id} {'enabled' if request.enabled else 'disabled'} by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return AutomationToggleResponse(
|
||||
success=True,
|
||||
automation_id=automation_id,
|
||||
enabled=request.enabled,
|
||||
message=f"Automation {'enabled' if request.enabled else 'disabled'}"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to toggle automation {automation_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
# ====================================================================
|
||||
# History
|
||||
# ====================================================================
|
||||
|
||||
@router.get(
|
||||
"/history",
|
||||
response_model=HistoryResponse,
|
||||
responses={400: {"model": ErrorResponse}}
|
||||
)
|
||||
async def get_history(
|
||||
entity_id: str = Query(..., description="Entity ID to get history for"),
|
||||
hours: int = Query(24, ge=1, le=168, description="Hours of history (1-168)")
|
||||
):
|
||||
"""
|
||||
Get state history for a device
|
||||
|
||||
Query Parameters:
|
||||
- entity_id: Device entity ID (required)
|
||||
- hours: Number of hours of history (default 24, max 168/1 week)
|
||||
"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
history_data = await ha.get_history(entity_id, hours)
|
||||
|
||||
# Transform HA history format to our format
|
||||
history_entries = []
|
||||
if history_data and len(history_data) > 0:
|
||||
for entry in history_data[0]: # First array is our entity
|
||||
history_entries.append(HistoryEntry(
|
||||
state=entry.get("state", "unknown"),
|
||||
timestamp=entry.get("last_changed", ""),
|
||||
attributes=entry.get("attributes", {})
|
||||
))
|
||||
|
||||
return HistoryResponse(
|
||||
entity_id=entity_id,
|
||||
history=history_entries
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get history for {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
housekeeping_controller = HousekeepingController()
|
||||
@@ -4,14 +4,14 @@ Infrastructure Management Controller
|
||||
Provides API endpoints for automated infrastructure management,
|
||||
including service deployment, configuration, and monitoring setup.
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, Depends
|
||||
from fastapi import APIRouter, HTTPException, Depends, Request
|
||||
from fastapi.responses import PlainTextResponse
|
||||
from typing import List, Dict, Any, Optional, Union
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
from src.controllers.base import BaseController
|
||||
from src.clients.portainer_client import get_portainer_client
|
||||
from src.clients.npm_client import get_npm_client
|
||||
from src.clients.kuma_client import get_kuma_client
|
||||
from src.logging_config import get_logger
|
||||
from src import service_groups
|
||||
from src.auth.oidc import get_admin_user, get_forward_auth_admin
|
||||
@@ -963,7 +963,7 @@ class InfrastructureController(BaseController):
|
||||
"/services/{name}/stop",
|
||||
response_model=OperationResult,
|
||||
summary="Stop a service or service group",
|
||||
description="Stop a service or service group by pausing monitors and stopping containers. Requires admin authentication when accessed externally via api.schweitz.net."
|
||||
description="Stop a service or service group by stopping containers. Requires admin authentication when accessed externally via api.schweitz.net."
|
||||
)
|
||||
async def stop_service(
|
||||
name: str,
|
||||
@@ -974,8 +974,7 @@ class InfrastructureController(BaseController):
|
||||
|
||||
This will:
|
||||
1. Validate service can be stopped (not always-on)
|
||||
2. Pause Uptime Kuma monitors for all services in group
|
||||
3. Stop the Portainer stack(s)
|
||||
2. Stop the Portainer stack containers
|
||||
|
||||
Args:
|
||||
name: Service or group name
|
||||
@@ -984,7 +983,6 @@ class InfrastructureController(BaseController):
|
||||
Operation result with details
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
kuma = get_kuma_client()
|
||||
|
||||
try:
|
||||
# Get all services in the group
|
||||
@@ -997,24 +995,13 @@ class InfrastructureController(BaseController):
|
||||
|
||||
results = {
|
||||
"stopped_services": [],
|
||||
"paused_monitors": [],
|
||||
"errors": []
|
||||
}
|
||||
|
||||
# Stop each service
|
||||
for service_name in services:
|
||||
try:
|
||||
# 1. Pause Uptime Kuma monitor
|
||||
try:
|
||||
monitor_paused = await kuma.pause_monitor_by_name(service_name)
|
||||
if monitor_paused:
|
||||
results["paused_monitors"].append(service_name)
|
||||
logger.info(f"Paused Kuma monitor for {service_name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to pause Kuma monitor for {service_name}: {e}")
|
||||
results["errors"].append(f"Kuma pause failed for {service_name}: {str(e)}")
|
||||
|
||||
# 2. Stop Portainer stack
|
||||
# Stop Portainer stack containers
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == service_name.lower()),
|
||||
@@ -1025,9 +1012,6 @@ class InfrastructureController(BaseController):
|
||||
stack_id = stack.get("Id")
|
||||
endpoint_id = stack.get("EndpointId")
|
||||
|
||||
# Stop stack by deleting it (Portainer doesn't have a "stop" operation)
|
||||
# Note: This is destructive. For a gentler approach, we'd need to use docker compose stop
|
||||
# Let's use docker API instead
|
||||
logger.info(f"Stopping containers for stack: {service_name}")
|
||||
|
||||
# Get containers for this stack
|
||||
@@ -1078,7 +1062,7 @@ class InfrastructureController(BaseController):
|
||||
"/services/{name}/start",
|
||||
response_model=OperationResult,
|
||||
summary="Start a service or service group",
|
||||
description="Start a service or service group by starting containers and resuming monitors. Requires admin authentication when accessed externally via api.schweitz.net."
|
||||
description="Start a service or service group by starting containers. Requires admin authentication when accessed externally via api.schweitz.net."
|
||||
)
|
||||
async def start_service(
|
||||
name: str,
|
||||
@@ -1088,8 +1072,7 @@ class InfrastructureController(BaseController):
|
||||
Start a service or service group
|
||||
|
||||
This will:
|
||||
1. Start the Portainer stack(s)
|
||||
2. Resume Uptime Kuma monitors for all services in group
|
||||
1. Start the Portainer stack containers
|
||||
|
||||
Args:
|
||||
name: Service or group name
|
||||
@@ -1098,7 +1081,6 @@ class InfrastructureController(BaseController):
|
||||
Operation result with details
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
kuma = get_kuma_client()
|
||||
|
||||
try:
|
||||
# Get all services in the group
|
||||
@@ -1106,14 +1088,13 @@ class InfrastructureController(BaseController):
|
||||
|
||||
results = {
|
||||
"started_services": [],
|
||||
"resumed_monitors": [],
|
||||
"errors": []
|
||||
}
|
||||
|
||||
# Start each service
|
||||
for service_name in services:
|
||||
try:
|
||||
# 1. Start Portainer stack (start containers)
|
||||
# Start Portainer stack containers
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == service_name.lower()),
|
||||
@@ -1147,16 +1128,6 @@ class InfrastructureController(BaseController):
|
||||
})
|
||||
logger.info(f"Started service: {service_name}")
|
||||
|
||||
# 2. Resume Uptime Kuma monitor
|
||||
try:
|
||||
monitor_resumed = await kuma.resume_monitor_by_name(service_name)
|
||||
if monitor_resumed:
|
||||
results["resumed_monitors"].append(service_name)
|
||||
logger.info(f"Resumed Kuma monitor for {service_name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to resume Kuma monitor for {service_name}: {e}")
|
||||
results["errors"].append(f"Kuma resume failed for {service_name}: {str(e)}")
|
||||
|
||||
else:
|
||||
results["errors"].append(f"Stack not found: {service_name}")
|
||||
|
||||
@@ -1181,8 +1152,6 @@ class InfrastructureController(BaseController):
|
||||
logger.error(f"Failed to start service group '{name}': {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
# ===== Monitoring Endpoints =====
|
||||
|
||||
@router.get(
|
||||
"/widget-data",
|
||||
summary="Get combined data for service control widget",
|
||||
@@ -1190,11 +1159,10 @@ class InfrastructureController(BaseController):
|
||||
)
|
||||
async def get_widget_data():
|
||||
"""
|
||||
Get combined service and monitor data for the widget
|
||||
Get combined service data for the widget
|
||||
|
||||
Returns all data needed by service-control widget in a single call:
|
||||
- Service list with status and container counts
|
||||
- Monitor list with uptime percentages
|
||||
- Service groups and always-on list
|
||||
|
||||
This endpoint is designed for browser-based widgets to avoid
|
||||
@@ -1202,7 +1170,6 @@ class InfrastructureController(BaseController):
|
||||
"""
|
||||
try:
|
||||
portainer = get_portainer_client()
|
||||
kuma = get_kuma_client()
|
||||
npm = get_npm_client()
|
||||
|
||||
# Fetch services (same logic as /services endpoint)
|
||||
@@ -1252,38 +1219,9 @@ class InfrastructureController(BaseController):
|
||||
"containers_total": containers_total
|
||||
})
|
||||
|
||||
# Fetch monitors with real-time status from metrics endpoint
|
||||
monitors_list = []
|
||||
try:
|
||||
# Get real-time status from Prometheus metrics
|
||||
metrics_data = await kuma.get_metrics_status()
|
||||
|
||||
for monitor_name, monitor_info in metrics_data.items():
|
||||
# Status: 1=UP, 0=DOWN, 2=PENDING, 3=MAINTENANCE
|
||||
status = monitor_info.get('status', 0)
|
||||
|
||||
# Convert status to simple up/down for widget
|
||||
# Treat UP (1) as 100%, anything else as 0%
|
||||
status_percentage = 100.0 if status == 1 else 0.0
|
||||
|
||||
monitors_list.append({
|
||||
"id": None, # Not available from metrics
|
||||
"name": monitor_name,
|
||||
"uptime_24h": status_percentage, # Current status as percentage
|
||||
"active": True, # Assume active if in metrics
|
||||
"status": status, # 1=UP, 0=DOWN, 2=PENDING, 3=MAINTENANCE
|
||||
"response_time": monitor_info.get('response_time', 0)
|
||||
})
|
||||
|
||||
logger.info(f"Fetched status for {len(monitors_list)} monitors from metrics")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to fetch monitors: {e}")
|
||||
# Continue without monitor data rather than failing
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"services": services,
|
||||
"monitors": monitors_list,
|
||||
"service_groups": {
|
||||
"groups": service_groups.list_service_groups(),
|
||||
"always_on": list(service_groups.ALWAYS_ON_SERVICES),
|
||||
@@ -1295,169 +1233,6 @@ class InfrastructureController(BaseController):
|
||||
logger.error(f"Failed to fetch widget data: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to fetch widget data: {str(e)}")
|
||||
|
||||
@router.get(
|
||||
"/monitors",
|
||||
summary="List all monitors",
|
||||
response_model=Dict[str, Any]
|
||||
)
|
||||
async def list_monitors():
|
||||
"""
|
||||
List all Uptime Kuma monitors
|
||||
|
||||
Returns:
|
||||
List of monitors with their configurations
|
||||
"""
|
||||
try:
|
||||
kuma = get_kuma_client()
|
||||
monitors = await kuma.get_monitors()
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"monitors": monitors,
|
||||
"total": len(monitors)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list monitors: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to list monitors: {str(e)}")
|
||||
|
||||
@router.post(
|
||||
"/monitors",
|
||||
summary="Create a new monitor",
|
||||
description="Create a new Uptime Kuma monitor. Requires admin authentication.",
|
||||
response_model=Dict[str, Any]
|
||||
)
|
||||
async def create_monitor(
|
||||
monitor_config: Dict[str, Any],
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Create a new Uptime Kuma monitor
|
||||
|
||||
Args:
|
||||
monitor_config: Monitor configuration (name, type, hostname, port, etc.)
|
||||
|
||||
Returns:
|
||||
Created monitor details including ID
|
||||
"""
|
||||
try:
|
||||
kuma = get_kuma_client()
|
||||
created_monitor = await kuma.add_monitor(monitor_config)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Monitor '{monitor_config.get('name')}' created successfully",
|
||||
"monitor": created_monitor
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create monitor: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to create monitor: {str(e)}")
|
||||
|
||||
@router.get(
|
||||
"/monitors/{monitor_id}",
|
||||
summary="Get monitor details",
|
||||
response_model=Dict[str, Any]
|
||||
)
|
||||
async def get_monitor(monitor_id: int):
|
||||
"""
|
||||
Get details of a specific monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
Monitor configuration and status
|
||||
"""
|
||||
try:
|
||||
kuma = get_kuma_client()
|
||||
monitor = await kuma.get_monitor(monitor_id)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"monitor": monitor
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get monitor {monitor_id}: {e}")
|
||||
raise HTTPException(status_code=404, detail=f"Monitor {monitor_id} not found: {str(e)}")
|
||||
|
||||
@router.put(
|
||||
"/monitors/{monitor_id}",
|
||||
summary="Update a monitor",
|
||||
description="Update an existing Uptime Kuma monitor. Requires admin authentication.",
|
||||
response_model=Dict[str, Any]
|
||||
)
|
||||
async def update_monitor(
|
||||
monitor_id: int,
|
||||
updates: Dict[str, Any],
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Update an existing monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
updates: Fields to update
|
||||
|
||||
Returns:
|
||||
Updated monitor details
|
||||
"""
|
||||
try:
|
||||
kuma = get_kuma_client()
|
||||
|
||||
# Get existing monitor
|
||||
existing = await kuma.get_monitor(monitor_id)
|
||||
|
||||
# Merge updates
|
||||
monitor_config = existing.copy()
|
||||
monitor_config.update(updates)
|
||||
|
||||
# Update monitor
|
||||
updated_monitor = await kuma.update_monitor(monitor_id, monitor_config)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Monitor {monitor_id} updated successfully",
|
||||
"monitor": updated_monitor
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update monitor {monitor_id}: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to update monitor: {str(e)}")
|
||||
|
||||
@router.delete(
|
||||
"/monitors/{monitor_id}",
|
||||
summary="Delete a monitor",
|
||||
description="Delete an Uptime Kuma monitor. Requires admin authentication.",
|
||||
response_model=Dict[str, Any]
|
||||
)
|
||||
async def delete_monitor(
|
||||
monitor_id: int,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Delete a monitor
|
||||
|
||||
Args:
|
||||
monitor_id: Monitor identifier
|
||||
|
||||
Returns:
|
||||
Success confirmation
|
||||
"""
|
||||
try:
|
||||
kuma = get_kuma_client()
|
||||
await kuma.delete_monitor(monitor_id)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Monitor {monitor_id} deleted successfully"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to delete monitor {monitor_id}: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to delete monitor: {str(e)}")
|
||||
|
||||
# ========================================================================
|
||||
# Container Management Endpoints (for core-ai infrastructure tools)
|
||||
# ========================================================================
|
||||
@@ -1926,6 +1701,357 @@ class InfrastructureController(BaseController):
|
||||
logger.error(f"Failed to get container resources: {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
# ========================================================================
|
||||
# Container Delete Endpoint (for Tatlock Control Room)
|
||||
# ========================================================================
|
||||
|
||||
@router.delete(
|
||||
"/containers/{container_id}",
|
||||
status_code=204,
|
||||
summary="Delete a container"
|
||||
)
|
||||
async def delete_container(
|
||||
container_id: str,
|
||||
force: bool = False
|
||||
):
|
||||
"""
|
||||
Delete a Docker container.
|
||||
|
||||
Args:
|
||||
container_id: Full container ID
|
||||
force: Force remove running container (default: false)
|
||||
|
||||
Returns:
|
||||
204 No Content on success
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Deleting container '{container_id}' (force={force})")
|
||||
|
||||
try:
|
||||
# Get endpoint
|
||||
endpoints = await portainer.get_endpoints()
|
||||
if not endpoints:
|
||||
raise HTTPException(status_code=500, detail="No Portainer endpoints available")
|
||||
|
||||
endpoint_id = endpoints[0]['Id']
|
||||
|
||||
# Delete the container
|
||||
await portainer.delete_container(endpoint_id, container_id, force=force)
|
||||
logger.info(f"Successfully deleted container '{container_id}'")
|
||||
|
||||
return None # 204 No Content
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
if "404" in error_str or "no such container" in error_str:
|
||||
raise HTTPException(status_code=404, detail=f"Container '{container_id}' not found")
|
||||
elif "409" in error_str or "conflict" in error_str:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Cannot delete container '{container_id}': container is running. Use force=true to force remove."
|
||||
)
|
||||
else:
|
||||
logger.error(f"Failed to delete container '{container_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
# ========================================================================
|
||||
# Stack Management Endpoints (for Tatlock Control Room)
|
||||
# ========================================================================
|
||||
|
||||
@router.get(
|
||||
"/stacks/{stack_id}/compose",
|
||||
summary="Get stack compose YAML",
|
||||
response_class=PlainTextResponse
|
||||
)
|
||||
async def get_stack_compose(stack_id: str):
|
||||
"""
|
||||
Get the Docker Compose YAML for a stack.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
|
||||
Returns:
|
||||
Docker Compose YAML as text/yaml
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Getting compose file for stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
# Get compose file content
|
||||
compose_content = await portainer.get_stack_file(stack["Id"])
|
||||
|
||||
return PlainTextResponse(
|
||||
content=compose_content,
|
||||
media_type="text/yaml"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get compose for stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.put(
|
||||
"/stacks/{stack_id}/compose",
|
||||
status_code=204,
|
||||
summary="Update stack compose YAML"
|
||||
)
|
||||
async def update_stack_compose(stack_id: str, request: Request):
|
||||
"""
|
||||
Update the Docker Compose YAML for a stack.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
request: Raw YAML content in request body
|
||||
|
||||
Returns:
|
||||
204 No Content on success
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Updating compose file for stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Read raw YAML from request body
|
||||
compose_content = (await request.body()).decode("utf-8")
|
||||
|
||||
if not compose_content.strip():
|
||||
raise HTTPException(status_code=400, detail="Empty compose content")
|
||||
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
stack_int_id = stack["Id"]
|
||||
endpoint_id = stack.get("EndpointId")
|
||||
|
||||
# Update stack with new compose content
|
||||
await portainer.update_stack(
|
||||
stack_id=stack_int_id,
|
||||
stack_file_content=compose_content,
|
||||
endpoint_id=endpoint_id,
|
||||
prune=False,
|
||||
pull_image=False
|
||||
)
|
||||
|
||||
logger.info(f"Successfully updated compose for stack '{stack_id}'")
|
||||
return None # 204 No Content
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update compose for stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/stacks/{stack_id}/env",
|
||||
response_model=Dict[str, str],
|
||||
summary="Get stack environment variables"
|
||||
)
|
||||
async def get_stack_env(stack_id: str):
|
||||
"""
|
||||
Get environment variables for a stack.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
|
||||
Returns:
|
||||
Dictionary of environment variable name-value pairs
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Getting env vars for stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
# Get full stack details including env vars
|
||||
stack_details = await portainer.get_stack(stack["Id"])
|
||||
env_list = stack_details.get("Env", [])
|
||||
|
||||
# Convert from [{name, value}] to {name: value}
|
||||
env_dict = {item["name"]: item["value"] for item in env_list}
|
||||
|
||||
return env_dict
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get env for stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.put(
|
||||
"/stacks/{stack_id}/env",
|
||||
status_code=204,
|
||||
summary="Update stack environment variables"
|
||||
)
|
||||
async def update_stack_env(stack_id: str, env_vars: Dict[str, str]):
|
||||
"""
|
||||
Update environment variables for a stack.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
env_vars: Dictionary of environment variable name-value pairs
|
||||
|
||||
Returns:
|
||||
204 No Content on success
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Updating env vars for stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
stack_int_id = stack["Id"]
|
||||
endpoint_id = stack.get("EndpointId")
|
||||
|
||||
# Convert from {name: value} to [{name, value}]
|
||||
env_list = [{"name": k, "value": v} for k, v in env_vars.items()]
|
||||
|
||||
# Update stack env vars
|
||||
await portainer.update_stack_env(
|
||||
stack_id=stack_int_id,
|
||||
endpoint_id=endpoint_id,
|
||||
env_vars=env_list
|
||||
)
|
||||
|
||||
logger.info(f"Successfully updated env vars for stack '{stack_id}'")
|
||||
return None # 204 No Content
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update env for stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.post(
|
||||
"/stacks/{stack_id}/deploy",
|
||||
status_code=202,
|
||||
summary="Deploy stack"
|
||||
)
|
||||
async def deploy_stack(stack_id: str):
|
||||
"""
|
||||
Redeploy a stack with current YAML and environment variables.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
|
||||
Returns:
|
||||
202 Accepted
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Deploying stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
stack_int_id = stack["Id"]
|
||||
endpoint_id = stack.get("EndpointId")
|
||||
|
||||
# Redeploy without pulling new images
|
||||
await portainer.redeploy_stack(
|
||||
stack_id=stack_int_id,
|
||||
endpoint_id=endpoint_id,
|
||||
pull_image=False
|
||||
)
|
||||
|
||||
logger.info(f"Successfully deployed stack '{stack_id}'")
|
||||
return None # 202 Accepted
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to deploy stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.post(
|
||||
"/stacks/{stack_id}/rebuild",
|
||||
status_code=202,
|
||||
summary="Rebuild stack"
|
||||
)
|
||||
async def rebuild_stack(stack_id: str):
|
||||
"""
|
||||
Pull fresh images and recreate all containers in the stack.
|
||||
|
||||
Args:
|
||||
stack_id: Stack name (compose project name)
|
||||
|
||||
Returns:
|
||||
202 Accepted
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
logger.info(f"Rebuilding stack '{stack_id}'")
|
||||
|
||||
try:
|
||||
# Find stack by name
|
||||
stacks = await portainer.get_stacks()
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == stack_id.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Stack '{stack_id}' not found")
|
||||
|
||||
stack_int_id = stack["Id"]
|
||||
endpoint_id = stack.get("EndpointId")
|
||||
|
||||
# Redeploy with image pull
|
||||
await portainer.redeploy_stack(
|
||||
stack_id=stack_int_id,
|
||||
endpoint_id=endpoint_id,
|
||||
pull_image=True
|
||||
)
|
||||
|
||||
logger.info(f"Successfully rebuilt stack '{stack_id}'")
|
||||
return None # 202 Accepted
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to rebuild stack '{stack_id}': {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
return router
|
||||
|
||||
|
||||
|
||||
@@ -2,16 +2,12 @@
|
||||
Tools Controller
|
||||
|
||||
Provides utility tool endpoints including:
|
||||
- Web scraping and content extraction
|
||||
- DNS lookups
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from src.controllers.base import BaseController
|
||||
from src.logging_config import get_logger
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.service import WebScraperService
|
||||
from src.web_scraper.exceptions import FetchError, ScrapingError
|
||||
from src.dns.schemas import DNSLookupRequest, DNSLookupResponse
|
||||
from src.dns.service import DNSService
|
||||
from src.dns.exceptions import DNSQueryError
|
||||
@@ -24,79 +20,17 @@ class ToolsController(BaseController):
|
||||
Controller for utility tools
|
||||
|
||||
Provides endpoints for:
|
||||
- Web scraping and content extraction
|
||||
- DNS lookups
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/tools", tags=["Tools"])
|
||||
# Initialize services (could be dependency injected for testing)
|
||||
self.scraper_service = WebScraperService()
|
||||
self.dns_service = DNSService()
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.post(
|
||||
"/scrape",
|
||||
response_model=WebScraperResponse,
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Scrape website content",
|
||||
description="""
|
||||
Scrape and extract main content from a website.
|
||||
|
||||
Uses trafilatura for intelligent content extraction (articles, blog posts, documentation),
|
||||
with BeautifulSoup as fallback. Perfect for feeding webpage content to LLMs.
|
||||
|
||||
**Features:**
|
||||
- Intelligent main content extraction
|
||||
- Removes navigation, ads, footers
|
||||
- Optional link extraction
|
||||
- Configurable content length limits
|
||||
|
||||
**Rate Limiting:** None (internal network use only)
|
||||
"""
|
||||
)
|
||||
async def scrape_website(request: WebScraperRequest) -> WebScraperResponse:
|
||||
"""
|
||||
Scrape a website and extract its main content
|
||||
|
||||
Args:
|
||||
request: Scraping request with URL and options
|
||||
|
||||
Returns:
|
||||
Extracted content with metadata
|
||||
|
||||
Raises:
|
||||
HTTPException: 400 for fetch errors, 500 for processing errors
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Received scrape request for: {request.url}")
|
||||
result = await self.scraper_service.scrape_url(request)
|
||||
return result
|
||||
|
||||
except FetchError as e:
|
||||
logger.warning(f"Fetch failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Failed to fetch URL: {str(e)}"
|
||||
)
|
||||
|
||||
except ScrapingError as e:
|
||||
logger.error(f"Scraping failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to extract content: {str(e)}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error: {str(e)}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="An unexpected error occurred"
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/dns/lookup",
|
||||
response_model=DNSLookupResponse,
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
"""
|
||||
Infrastructure Credentials Template
|
||||
|
||||
INSTRUCTIONS:
|
||||
1. Copy this file to credentials.py
|
||||
2. Fill in your actual credentials
|
||||
3. DO NOT commit credentials.py to version control (it's in .gitignore)
|
||||
|
||||
This file should be committed to the repository as a template.
|
||||
"""
|
||||
|
||||
# Portainer Configuration
|
||||
PORTAINER_URL = "http://localhost:8001"
|
||||
PORTAINER_API_KEY = "ptr_your_api_token_here" # Create in Portainer UI: User menu → My account → Access tokens
|
||||
|
||||
# Nginx Proxy Manager Configuration
|
||||
NPM_URL = "http://localhost:81"
|
||||
NPM_EMAIL = "admin@example.com"
|
||||
NPM_PASSWORD = "your_password_here"
|
||||
|
||||
# Uptime Kuma Configuration
|
||||
KUMA_URL = "http://localhost:3001"
|
||||
KUMA_USERNAME = "admin"
|
||||
KUMA_PASSWORD = "your_password_here"
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
Database package for Core-API
|
||||
|
||||
Provides async PostgreSQL database connectivity using SQLAlchemy 2.0.
|
||||
"""
|
||||
from src.db.database import (
|
||||
get_async_session,
|
||||
get_database,
|
||||
Database,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"get_async_session",
|
||||
"get_database",
|
||||
"Database",
|
||||
]
|
||||
@@ -0,0 +1,179 @@
|
||||
"""
|
||||
Database Connection Module
|
||||
|
||||
Provides async PostgreSQL connectivity using SQLAlchemy 2.0 with asyncpg driver.
|
||||
Follows the existing singleton pattern used throughout core-api.
|
||||
"""
|
||||
from typing import AsyncGenerator, Optional
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncSession,
|
||||
AsyncEngine,
|
||||
create_async_engine,
|
||||
async_sessionmaker,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from src.config import get_settings
|
||||
from src.logging_config import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
"""
|
||||
SQLAlchemy declarative base for all models
|
||||
|
||||
All database models should inherit from this class.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class Database:
|
||||
"""
|
||||
Async database connection manager
|
||||
|
||||
Provides async engine and session factory for PostgreSQL connections.
|
||||
Uses asyncpg driver for optimal async performance.
|
||||
"""
|
||||
|
||||
def __init__(self, database_url: Optional[str] = None):
|
||||
"""
|
||||
Initialize database connection manager
|
||||
|
||||
Args:
|
||||
database_url: PostgreSQL connection URL (default from settings)
|
||||
"""
|
||||
# Convert postgresql:// to postgresql+asyncpg:// for async driver
|
||||
url = database_url or settings.database_url
|
||||
if url.startswith("postgresql://"):
|
||||
url = url.replace("postgresql://", "postgresql+asyncpg://", 1)
|
||||
|
||||
self._url = url
|
||||
self._engine: Optional[AsyncEngine] = None
|
||||
self._session_factory: Optional[async_sessionmaker[AsyncSession]] = None
|
||||
|
||||
@property
|
||||
def engine(self) -> AsyncEngine:
|
||||
"""
|
||||
Get or create the async database engine
|
||||
|
||||
Returns:
|
||||
AsyncEngine instance
|
||||
"""
|
||||
if self._engine is None:
|
||||
self._engine = create_async_engine(
|
||||
self._url,
|
||||
echo=settings.debug, # Log SQL in debug mode
|
||||
poolclass=NullPool, # Disable connection pooling for serverless compatibility
|
||||
)
|
||||
logger.info(f"Database engine created for {self._url.split('@')[-1]}")
|
||||
return self._engine
|
||||
|
||||
@property
|
||||
def session_factory(self) -> async_sessionmaker[AsyncSession]:
|
||||
"""
|
||||
Get or create the async session factory
|
||||
|
||||
Returns:
|
||||
Session factory for creating database sessions
|
||||
"""
|
||||
if self._session_factory is None:
|
||||
self._session_factory = async_sessionmaker(
|
||||
bind=self.engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
autocommit=False,
|
||||
autoflush=False,
|
||||
)
|
||||
return self._session_factory
|
||||
|
||||
async def create_tables(self) -> None:
|
||||
"""
|
||||
Create all database tables
|
||||
|
||||
Should only be used for development/testing.
|
||||
Use Alembic migrations for production.
|
||||
"""
|
||||
async with self.engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("Database tables created")
|
||||
|
||||
async def drop_tables(self) -> None:
|
||||
"""
|
||||
Drop all database tables
|
||||
|
||||
WARNING: Destroys all data. Use with caution.
|
||||
"""
|
||||
async with self.engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
logger.warning("Database tables dropped")
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if database connection is healthy
|
||||
|
||||
Returns:
|
||||
True if connection successful, False otherwise
|
||||
"""
|
||||
try:
|
||||
async with self.session_factory() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Database health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def close(self) -> None:
|
||||
"""
|
||||
Close database connections and dispose of engine
|
||||
"""
|
||||
if self._engine is not None:
|
||||
await self._engine.dispose()
|
||||
self._engine = None
|
||||
self._session_factory = None
|
||||
logger.info("Database connections closed")
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_database: Optional[Database] = None
|
||||
|
||||
|
||||
def get_database() -> Database:
|
||||
"""
|
||||
Get singleton database instance
|
||||
|
||||
Returns:
|
||||
Database instance
|
||||
"""
|
||||
global _database
|
||||
if _database is None:
|
||||
_database = Database()
|
||||
return _database
|
||||
|
||||
|
||||
async def get_async_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""
|
||||
FastAPI dependency for database sessions
|
||||
|
||||
Yields an async session that is automatically closed after the request.
|
||||
|
||||
Usage:
|
||||
@router.get("/items")
|
||||
async def get_items(session: AsyncSession = Depends(get_async_session)):
|
||||
result = await session.execute(select(Item))
|
||||
return result.scalars().all()
|
||||
|
||||
Yields:
|
||||
AsyncSession instance
|
||||
"""
|
||||
database = get_database()
|
||||
async with database.session_factory() as session:
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
@@ -0,0 +1,19 @@
|
||||
"""
|
||||
SQLAlchemy Models for Core-API
|
||||
|
||||
Database models for authentication, authorization, and user management.
|
||||
"""
|
||||
from src.db.models.user import User
|
||||
from src.db.models.role import Role, UserRole
|
||||
from src.db.models.user_preferences import UserPreferences
|
||||
from src.db.models.api_key import ApiKey
|
||||
from src.db.models.group import Group
|
||||
|
||||
__all__ = [
|
||||
"User",
|
||||
"Role",
|
||||
"UserRole",
|
||||
"UserPreferences",
|
||||
"ApiKey",
|
||||
"Group",
|
||||
]
|
||||
@@ -0,0 +1,98 @@
|
||||
"""
|
||||
API Key Model
|
||||
|
||||
Provides API key authentication as fallback for OIDC.
|
||||
Keys are tied to user accounts and inherit user permissions.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import String, ForeignKey, DateTime, func
|
||||
from sqlalchemy.dialects.postgresql import UUID, ARRAY
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.db.database import Base
|
||||
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.db.models.user import User
|
||||
|
||||
|
||||
class ApiKey(Base):
|
||||
"""
|
||||
API Key model for programmatic access
|
||||
|
||||
API keys provide an alternative to OIDC for:
|
||||
- Local development without SSO
|
||||
- Service-to-service communication
|
||||
- Scripts and automation
|
||||
|
||||
Keys inherit the user's roles but can optionally
|
||||
be restricted to a subset of scopes.
|
||||
"""
|
||||
|
||||
__tablename__ = "api_keys"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(100),
|
||||
nullable=False,
|
||||
comment="Human-readable key name (e.g., 'Dev Laptop', 'CI/CD')",
|
||||
)
|
||||
key_hash: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
nullable=False,
|
||||
comment="SHA-256 hash of the API key",
|
||||
)
|
||||
key_prefix: Mapped[str] = mapped_column(
|
||||
String(8),
|
||||
nullable=False,
|
||||
comment="First 8 chars of key for identification (e.g., 'cak_abc1')",
|
||||
)
|
||||
scopes: Mapped[List[str] | None] = mapped_column(
|
||||
ARRAY(String),
|
||||
nullable=True,
|
||||
comment="Optional scope restriction (subset of user roles)",
|
||||
)
|
||||
expires_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
comment="Optional expiration timestamp",
|
||||
)
|
||||
last_used_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
comment="Last time this key was used",
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User",
|
||||
back_populates="api_keys",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<ApiKey {self.key_prefix}... ({self.name})>"
|
||||
|
||||
@property
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if the API key has expired"""
|
||||
if self.expires_at is None:
|
||||
return False
|
||||
return datetime.now(self.expires_at.tzinfo) > self.expires_at
|
||||
@@ -0,0 +1,85 @@
|
||||
"""
|
||||
Group Model
|
||||
|
||||
Represents groups synced from Authentik.
|
||||
Groups are used for access control and user organization.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
from sqlalchemy import String, Boolean, DateTime, func, Table, Column, ForeignKey
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.db.database import Base
|
||||
|
||||
|
||||
# Association table for User-Group many-to-many relationship
|
||||
user_groups = Table(
|
||||
"user_groups",
|
||||
Base.metadata,
|
||||
Column("user_id", UUID(as_uuid=True), ForeignKey("users.id", ondelete="CASCADE"), primary_key=True),
|
||||
Column("group_id", UUID(as_uuid=True), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True),
|
||||
)
|
||||
|
||||
|
||||
class Group(Base):
|
||||
"""
|
||||
Group model synced from Authentik
|
||||
|
||||
Groups are fetched from Authentik admin API and cached locally.
|
||||
They represent organizational units for access control.
|
||||
"""
|
||||
|
||||
__tablename__ = "groups"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
authentik_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Authentik group UUID",
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
is_superuser: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=False,
|
||||
nullable=False,
|
||||
comment="Whether members of this group have superuser privileges",
|
||||
)
|
||||
parent_name: Mapped[str | None] = mapped_column(
|
||||
String(255),
|
||||
nullable=True,
|
||||
comment="Parent group name for hierarchy",
|
||||
)
|
||||
member_count: Mapped[int] = mapped_column(
|
||||
default=0,
|
||||
nullable=False,
|
||||
comment="Number of users in this group",
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
synced_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
comment="Last sync from Authentik",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Group {self.name}>"
|
||||
@@ -0,0 +1,93 @@
|
||||
"""
|
||||
Role Models
|
||||
|
||||
Defines domain-scoped permissions mapped from Authentik groups.
|
||||
Format: {domain}:{action} (e.g., control-room:admin, media:viewer)
|
||||
"""
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
from sqlalchemy import String, ForeignKey
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.db.database import Base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.db.models.user import User
|
||||
|
||||
|
||||
class Role(Base):
|
||||
"""
|
||||
Role model for domain-scoped permissions
|
||||
|
||||
Roles are seeded from configuration, not user-editable.
|
||||
Each role maps to an Authentik group (e.g., tatlock-control-room-admin).
|
||||
|
||||
Domains: control-room, library, media, ai, housekeeper, developer, documents, gaming, admin
|
||||
Actions: viewer, user, editor, admin (hierarchical)
|
||||
"""
|
||||
|
||||
__tablename__ = "roles"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(100),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Role name in format domain:action (e.g., control-room:admin)",
|
||||
)
|
||||
domain: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Permission domain (e.g., control-room, media, ai)",
|
||||
)
|
||||
action: Mapped[str] = mapped_column(
|
||||
String(20),
|
||||
nullable=False,
|
||||
comment="Permission action (viewer, user, editor, admin)",
|
||||
)
|
||||
authentik_group: Mapped[str | None] = mapped_column(
|
||||
String(255),
|
||||
nullable=True,
|
||||
unique=True,
|
||||
comment="Corresponding Authentik group name (e.g., tatlock-control-room-admin)",
|
||||
)
|
||||
|
||||
# Relationships
|
||||
users: Mapped[List["User"]] = relationship(
|
||||
"User",
|
||||
secondary="user_roles",
|
||||
back_populates="roles",
|
||||
lazy="selectin",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Role {self.name}>"
|
||||
|
||||
|
||||
class UserRole(Base):
|
||||
"""
|
||||
Association table for User-Role many-to-many relationship
|
||||
|
||||
Synced from Authentik groups during user authentication.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_roles"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
role_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("roles.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
@@ -0,0 +1,94 @@
|
||||
"""
|
||||
User Model
|
||||
|
||||
Represents users synced from Authentik SSO.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
from sqlalchemy import String, Boolean, DateTime, func
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.db.database import Base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.db.models.role import Role
|
||||
from src.db.models.user_preferences import UserPreferences
|
||||
from src.db.models.api_key import ApiKey
|
||||
|
||||
|
||||
class User(Base):
|
||||
"""
|
||||
User model synced from Authentik
|
||||
|
||||
Users are created/updated when they authenticate via OIDC.
|
||||
The authentik_id links to the Authentik user record.
|
||||
"""
|
||||
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
authentik_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
email: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
nullable=False,
|
||||
)
|
||||
avatar_url: Mapped[str | None] = mapped_column(
|
||||
String(500),
|
||||
nullable=True,
|
||||
)
|
||||
api_keys_enabled: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=True,
|
||||
nullable=False,
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
last_login: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
roles: Mapped[List["Role"]] = relationship(
|
||||
"Role",
|
||||
secondary="user_roles",
|
||||
back_populates="users",
|
||||
lazy="selectin",
|
||||
)
|
||||
preferences: Mapped["UserPreferences"] = relationship(
|
||||
"UserPreferences",
|
||||
back_populates="user",
|
||||
uselist=False,
|
||||
lazy="selectin",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
api_keys: Mapped[List["ApiKey"]] = relationship(
|
||||
"ApiKey",
|
||||
back_populates="user",
|
||||
lazy="selectin",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<User {self.email}>"
|
||||
@@ -0,0 +1,61 @@
|
||||
"""
|
||||
User Preferences Model
|
||||
|
||||
Stores user-specific settings like theme and default room.
|
||||
"""
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import String, ForeignKey
|
||||
from sqlalchemy.dialects.postgresql import UUID, JSONB
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.db.database import Base
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.db.models.user import User
|
||||
|
||||
|
||||
class UserPreferences(Base):
|
||||
"""
|
||||
User preferences model
|
||||
|
||||
Stores user-specific settings that persist across sessions.
|
||||
Extended settings stored in preferences_json for flexibility.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_preferences"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
theme: Mapped[str] = mapped_column(
|
||||
String(20),
|
||||
default="system",
|
||||
nullable=False,
|
||||
comment="Theme preference: system, light, dark",
|
||||
)
|
||||
default_room: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
default="front-hall",
|
||||
nullable=False,
|
||||
comment="Default room for housekeeping features",
|
||||
)
|
||||
preferences_json: Mapped[dict] = mapped_column(
|
||||
JSONB,
|
||||
default=dict,
|
||||
nullable=False,
|
||||
comment="Extended preferences as JSON",
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User",
|
||||
back_populates="preferences",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<UserPreferences user_id={self.user_id}>"
|
||||
@@ -0,0 +1,5 @@
|
||||
"""
|
||||
Domain modules for Core-API
|
||||
|
||||
Each domain contains its own models, schemas, services, and controllers.
|
||||
"""
|
||||
@@ -0,0 +1,62 @@
|
||||
"""
|
||||
Authentication Domain
|
||||
|
||||
Provides OIDC/OAuth2 authentication via Authentik, user management,
|
||||
roles, groups, and API key authentication.
|
||||
"""
|
||||
from src.domains.auth.oidc import (
|
||||
get_current_user,
|
||||
get_admin_user,
|
||||
get_optional_user,
|
||||
get_forward_auth_user,
|
||||
get_forward_auth_admin,
|
||||
oidc_config,
|
||||
# Permission system
|
||||
require_permission,
|
||||
require_any_permission,
|
||||
ACTION_HIERARCHY,
|
||||
VALID_DOMAINS,
|
||||
DEFAULT_CATEGORY,
|
||||
)
|
||||
from src.domains.auth.service import AuthService, get_auth_service
|
||||
from src.domains.auth.controller import auth_controller
|
||||
from src.domains.auth.models import (
|
||||
User,
|
||||
Role,
|
||||
UserRole,
|
||||
Group,
|
||||
UserPreferences,
|
||||
ApiKey,
|
||||
user_groups,
|
||||
group_roles,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# OIDC dependencies
|
||||
"get_current_user",
|
||||
"get_admin_user",
|
||||
"get_optional_user",
|
||||
"get_forward_auth_user",
|
||||
"get_forward_auth_admin",
|
||||
"oidc_config",
|
||||
# Permission system
|
||||
"require_permission",
|
||||
"require_any_permission",
|
||||
"ACTION_HIERARCHY",
|
||||
"VALID_DOMAINS",
|
||||
"DEFAULT_CATEGORY",
|
||||
# Service
|
||||
"AuthService",
|
||||
"get_auth_service",
|
||||
# Controller
|
||||
"auth_controller",
|
||||
# Models
|
||||
"User",
|
||||
"Role",
|
||||
"UserRole",
|
||||
"Group",
|
||||
"UserPreferences",
|
||||
"ApiKey",
|
||||
"user_groups",
|
||||
"group_roles",
|
||||
]
|
||||
@@ -0,0 +1,592 @@
|
||||
"""
|
||||
Authentication Controller
|
||||
|
||||
Provides authentication endpoints for OIDC token sync and user management.
|
||||
"""
|
||||
import uuid
|
||||
from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Path, Query
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.logging import get_logger
|
||||
from src.shared.database import get_async_session
|
||||
from src.domains.auth.schemas import (
|
||||
AuthSyncRequest, AuthSyncResponse, UsersListResponse,
|
||||
BulkSyncResultSchema, GroupsListResponse, RolesListResponse,
|
||||
GroupRoleAssignmentResponse, UserProfileResponse, PreferencesUpdateRequest,
|
||||
UserPreferencesSchema, ApiKeyCreateRequest, ApiKeyCreateResponse,
|
||||
ApiKeysListResponse,
|
||||
)
|
||||
from src.domains.auth.service import AuthService
|
||||
from src.domains.auth.oidc import get_current_user
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AuthController(BaseController):
|
||||
"""
|
||||
Controller for authentication operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Token synchronization (login)
|
||||
- User profile retrieval
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/auth", tags=["Authentication"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.post(
|
||||
"/sync",
|
||||
summary="Sync user from OIDC token",
|
||||
response_model=AuthSyncResponse,
|
||||
responses={
|
||||
200: {"description": "User synced successfully"},
|
||||
401: {"description": "Invalid or expired token"},
|
||||
503: {"description": "Authentication service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_user(
|
||||
request: AuthSyncRequest,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> AuthSyncResponse:
|
||||
"""
|
||||
Synchronize user from OIDC access token
|
||||
|
||||
This endpoint should be called after the client obtains an access token
|
||||
from Authentik. It:
|
||||
1. Validates the token via Authentik's userinfo endpoint
|
||||
2. Creates or updates the user in the database
|
||||
3. Syncs roles from Authentik groups
|
||||
4. Returns the user profile with roles and preferences
|
||||
|
||||
The client should store the returned user info for local use.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
# Validate token with Authentik
|
||||
token_info = await service.validate_token(request.access_token)
|
||||
except ValueError as e:
|
||||
logger.warning(f"Token validation failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
# Sync user to database
|
||||
user, is_new = await service.sync_user(token_info)
|
||||
|
||||
# Sync roles from groups
|
||||
roles = await service.sync_roles(user, token_info.groups)
|
||||
|
||||
# Commit the transaction
|
||||
await session.commit()
|
||||
|
||||
# Refresh to get relationships
|
||||
await session.refresh(user, ["preferences"])
|
||||
|
||||
# Build response
|
||||
return AuthSyncResponse(
|
||||
user=service.user_to_schema(user),
|
||||
roles=service.roles_to_schema(roles),
|
||||
preferences=service.preferences_to_schema(user.preferences),
|
||||
is_new_user=is_new,
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/users",
|
||||
summary="List all users",
|
||||
response_model=UsersListResponse,
|
||||
responses={
|
||||
200: {"description": "List of users"},
|
||||
},
|
||||
)
|
||||
async def list_users(
|
||||
search: Optional[str] = Query(None, description="Search by name or email"),
|
||||
offset: int = Query(0, ge=0, description="Number of records to skip"),
|
||||
limit: int = Query(50, ge=1, le=100, description="Maximum records to return"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> UsersListResponse:
|
||||
"""
|
||||
List all users who have logged in via Authentik
|
||||
|
||||
Returns paginated list of users with their roles.
|
||||
Supports search filtering by name or email.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
items, total = await service.list_users(
|
||||
search=search,
|
||||
offset=offset,
|
||||
limit=limit,
|
||||
)
|
||||
return UsersListResponse(items=items, total=total)
|
||||
|
||||
@router.post(
|
||||
"/users/sync-from-authentik",
|
||||
summary="Bulk sync users from Authentik",
|
||||
response_model=BulkSyncResultSchema,
|
||||
responses={
|
||||
200: {"description": "Sync completed"},
|
||||
401: {"description": "Authentik API token invalid"},
|
||||
503: {"description": "Authentik service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_users_from_authentik(
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all users from Authentik and sync to local database
|
||||
|
||||
This endpoint uses the Authentik admin API to fetch all users
|
||||
and create/update them in the local database. Requires
|
||||
AUTHENTIK_CORE_API_TOKEN to be configured.
|
||||
|
||||
Use this to initially populate users or to re-sync after
|
||||
changes in Authentik.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
result = await service.bulk_sync_from_authentik()
|
||||
logger.info(
|
||||
f"Bulk sync completed: {result.created} created, "
|
||||
f"{result.updated} updated, {result.failed} failed"
|
||||
)
|
||||
return result
|
||||
except ValueError as e:
|
||||
logger.error(f"Bulk sync failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/groups",
|
||||
summary="List all groups",
|
||||
response_model=GroupsListResponse,
|
||||
responses={
|
||||
200: {"description": "List of groups"},
|
||||
},
|
||||
)
|
||||
async def list_groups(
|
||||
search: Optional[str] = Query(None, description="Search by group name"),
|
||||
offset: int = Query(0, ge=0, description="Number of records to skip"),
|
||||
limit: int = Query(50, ge=1, le=100, description="Maximum records to return"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> GroupsListResponse:
|
||||
"""
|
||||
List all groups synced from Authentik
|
||||
|
||||
Returns paginated list of groups with their details.
|
||||
Supports search filtering by name.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
items, total = await service.list_groups(
|
||||
search=search,
|
||||
offset=offset,
|
||||
limit=limit,
|
||||
)
|
||||
return GroupsListResponse(items=items, total=total)
|
||||
|
||||
@router.post(
|
||||
"/groups/sync-from-authentik",
|
||||
summary="Bulk sync groups from Authentik",
|
||||
response_model=BulkSyncResultSchema,
|
||||
responses={
|
||||
200: {"description": "Sync completed"},
|
||||
401: {"description": "Authentik API credentials invalid"},
|
||||
503: {"description": "Authentik service unavailable"},
|
||||
},
|
||||
)
|
||||
async def sync_groups_from_authentik(
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> BulkSyncResultSchema:
|
||||
"""
|
||||
Fetch all groups from Authentik and sync to local database
|
||||
|
||||
This endpoint uses the Authentik admin API to fetch all groups
|
||||
and create/update them in the local database. Requires
|
||||
AUTHENTIK_USERNAME and AUTHENTIK_PASSWORD to be configured.
|
||||
|
||||
Use this to populate groups or to re-sync after changes in Authentik.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
result = await service.bulk_sync_groups_from_authentik()
|
||||
logger.info(
|
||||
f"Groups bulk sync completed: {result.created} created, "
|
||||
f"{result.updated} updated, {result.failed} failed"
|
||||
)
|
||||
return result
|
||||
except ValueError as e:
|
||||
logger.error(f"Groups bulk sync failed: {e}")
|
||||
raise HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/roles",
|
||||
summary="List all roles",
|
||||
response_model=RolesListResponse,
|
||||
responses={
|
||||
200: {"description": "List of all available roles"},
|
||||
},
|
||||
)
|
||||
async def list_roles(
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> RolesListResponse:
|
||||
"""
|
||||
List all available roles in the system
|
||||
|
||||
Returns all domain.category:action role combinations.
|
||||
Use these when assigning roles to groups.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
roles = await service.list_roles()
|
||||
return RolesListResponse(
|
||||
items=service.roles_to_schema(roles),
|
||||
total=len(roles),
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/groups/{group_id}/roles/{role_id}",
|
||||
summary="Assign role to group",
|
||||
response_model=GroupRoleAssignmentResponse,
|
||||
responses={
|
||||
200: {"description": "Role assigned successfully"},
|
||||
404: {"description": "Group or role not found"},
|
||||
},
|
||||
)
|
||||
async def assign_role_to_group(
|
||||
group_id: uuid.UUID = Path(..., description="Group ID"),
|
||||
role_id: uuid.UUID = Path(..., description="Role ID to assign"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> GroupRoleAssignmentResponse:
|
||||
"""
|
||||
Assign a role to a group
|
||||
|
||||
All users in this group will inherit this role's permissions.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
group = await service.assign_role_to_group(group_id, role_id)
|
||||
await session.commit()
|
||||
return GroupRoleAssignmentResponse(
|
||||
group_id=group.id,
|
||||
group_name=group.name,
|
||||
roles=[role.name for role in group.roles],
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
|
||||
@router.delete(
|
||||
"/groups/{group_id}/roles/{role_id}",
|
||||
summary="Remove role from group",
|
||||
response_model=GroupRoleAssignmentResponse,
|
||||
responses={
|
||||
200: {"description": "Role removed successfully"},
|
||||
404: {"description": "Group or role not found"},
|
||||
},
|
||||
)
|
||||
async def remove_role_from_group(
|
||||
group_id: uuid.UUID = Path(..., description="Group ID"),
|
||||
role_id: uuid.UUID = Path(..., description="Role ID to remove"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> GroupRoleAssignmentResponse:
|
||||
"""
|
||||
Remove a role from a group
|
||||
|
||||
Users in this group will no longer inherit this role's permissions.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
try:
|
||||
group = await service.remove_role_from_group(group_id, role_id)
|
||||
await session.commit()
|
||||
return GroupRoleAssignmentResponse(
|
||||
group_id=group.id,
|
||||
group_name=group.name,
|
||||
roles=[role.name for role in group.roles],
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
|
||||
# =====================================================================
|
||||
# Phase 4: User Profile & Settings
|
||||
# =====================================================================
|
||||
|
||||
@router.get(
|
||||
"/users/me",
|
||||
summary="Get current user profile",
|
||||
response_model=UserProfileResponse,
|
||||
responses={
|
||||
200: {"description": "User profile with roles and preferences"},
|
||||
401: {"description": "Not authenticated"},
|
||||
404: {"description": "User not found in database"},
|
||||
},
|
||||
)
|
||||
async def get_current_user_profile(
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> UserProfileResponse:
|
||||
"""
|
||||
Get the current authenticated user's profile
|
||||
|
||||
Returns the user's profile, roles, and preferences.
|
||||
Requires authentication via Bearer token or API key.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
# Get authentik_id from claims (JWT 'sub' field)
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="User not found - please sync via /auth/sync first",
|
||||
)
|
||||
|
||||
return UserProfileResponse(
|
||||
user=service.user_to_schema(user),
|
||||
roles=service.roles_to_schema(user.roles),
|
||||
preferences=service.preferences_to_schema(user.preferences),
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/users/me/preferences",
|
||||
summary="Get user preferences",
|
||||
response_model=UserPreferencesSchema,
|
||||
responses={
|
||||
200: {"description": "User preferences"},
|
||||
401: {"description": "Not authenticated"},
|
||||
},
|
||||
)
|
||||
async def get_preferences(
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> UserPreferencesSchema:
|
||||
"""
|
||||
Get the current user's preferences
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
return service.preferences_to_schema(user.preferences)
|
||||
|
||||
@router.patch(
|
||||
"/users/me/preferences",
|
||||
summary="Update user preferences",
|
||||
response_model=UserPreferencesSchema,
|
||||
responses={
|
||||
200: {"description": "Updated preferences"},
|
||||
401: {"description": "Not authenticated"},
|
||||
422: {"description": "Invalid preference value"},
|
||||
},
|
||||
)
|
||||
async def update_preferences(
|
||||
request: PreferencesUpdateRequest,
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> UserPreferencesSchema:
|
||||
"""
|
||||
Update the current user's preferences
|
||||
|
||||
Only provided fields are updated. preferences_json is merged
|
||||
with existing values (not replaced).
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
try:
|
||||
prefs = await service.update_preferences(
|
||||
user_id=user.id,
|
||||
theme=request.theme,
|
||||
default_room=request.default_room,
|
||||
preferences_json=request.preferences_json,
|
||||
)
|
||||
await session.commit()
|
||||
return service.preferences_to_schema(prefs)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=422, detail=str(e))
|
||||
|
||||
# =====================================================================
|
||||
# Phase 4: API Keys
|
||||
# =====================================================================
|
||||
|
||||
@router.get(
|
||||
"/users/me/api-keys",
|
||||
summary="List user's API keys",
|
||||
response_model=ApiKeysListResponse,
|
||||
responses={
|
||||
200: {"description": "List of API keys"},
|
||||
401: {"description": "Not authenticated"},
|
||||
},
|
||||
)
|
||||
async def list_api_keys(
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> ApiKeysListResponse:
|
||||
"""
|
||||
List all API keys for the current user
|
||||
|
||||
Returns key metadata only - the actual key values are never
|
||||
retrievable after creation.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
keys = await service.list_user_api_keys(user.id)
|
||||
return ApiKeysListResponse(
|
||||
items=[service.api_key_to_schema(k) for k in keys],
|
||||
total=len(keys),
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/users/me/api-keys",
|
||||
summary="Create a new API key",
|
||||
response_model=ApiKeyCreateResponse,
|
||||
responses={
|
||||
201: {"description": "API key created"},
|
||||
401: {"description": "Not authenticated"},
|
||||
403: {"description": "API keys disabled for user"},
|
||||
},
|
||||
)
|
||||
async def create_api_key(
|
||||
request: ApiKeyCreateRequest,
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> ApiKeyCreateResponse:
|
||||
"""
|
||||
Create a new API key for the current user
|
||||
|
||||
**IMPORTANT**: The full API key is only returned once in this response!
|
||||
Store it securely - it cannot be retrieved again.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
try:
|
||||
api_key, full_key = await service.create_api_key(
|
||||
user_id=user.id,
|
||||
name=request.name,
|
||||
scopes=request.scopes,
|
||||
expires_in_days=request.expires_in_days,
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
return ApiKeyCreateResponse(
|
||||
id=api_key.id,
|
||||
name=api_key.name,
|
||||
key=full_key, # Only time this is returned!
|
||||
key_prefix=api_key.key_prefix,
|
||||
scopes=api_key.scopes,
|
||||
expires_at=api_key.expires_at,
|
||||
created_at=api_key.created_at,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=403, detail=str(e))
|
||||
|
||||
@router.delete(
|
||||
"/users/me/api-keys/{key_id}",
|
||||
summary="Delete an API key",
|
||||
responses={
|
||||
204: {"description": "API key deleted"},
|
||||
401: {"description": "Not authenticated"},
|
||||
404: {"description": "API key not found"},
|
||||
},
|
||||
)
|
||||
async def delete_api_key(
|
||||
key_id: uuid.UUID = Path(..., description="API key ID to delete"),
|
||||
user_claims: dict = Depends(get_current_user),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
) -> JSONResponse:
|
||||
"""
|
||||
Delete an API key
|
||||
|
||||
The key will be immediately invalidated.
|
||||
"""
|
||||
service = AuthService(session)
|
||||
|
||||
authentik_id_str = user_claims.get("sub")
|
||||
if not authentik_id_str or authentik_id_str == "local-user":
|
||||
raise HTTPException(status_code=401, detail="Authentication required")
|
||||
|
||||
try:
|
||||
authentik_id = uuid.UUID(authentik_id_str)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=401, detail="Invalid user identifier")
|
||||
|
||||
user = await service.get_user_by_authentik_id(authentik_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
try:
|
||||
deleted = await service.delete_api_key(user.id, key_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=404, detail="API key not found")
|
||||
await session.commit()
|
||||
return JSONResponse(status_code=204, content=None)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=403, detail=str(e))
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
auth_controller = AuthController()
|
||||
@@ -0,0 +1,408 @@
|
||||
"""
|
||||
Authentication Domain Models
|
||||
|
||||
SQLAlchemy models for users, roles, groups, API keys, and preferences.
|
||||
All authentication-related database models consolidated in one file.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
from sqlalchemy import String, Boolean, DateTime, func, ForeignKey, Table, Column
|
||||
from sqlalchemy.dialects.postgresql import UUID, JSONB, ARRAY
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.shared.database import Base
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Association Tables
|
||||
# =============================================================================
|
||||
|
||||
user_groups = Table(
|
||||
"user_groups",
|
||||
Base.metadata,
|
||||
Column("user_id", UUID(as_uuid=True), ForeignKey("users.id", ondelete="CASCADE"), primary_key=True),
|
||||
Column("group_id", UUID(as_uuid=True), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True),
|
||||
)
|
||||
|
||||
group_roles = Table(
|
||||
"group_roles",
|
||||
Base.metadata,
|
||||
Column("group_id", UUID(as_uuid=True), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True),
|
||||
Column("role_id", UUID(as_uuid=True), ForeignKey("roles.id", ondelete="CASCADE"), primary_key=True),
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# User Model
|
||||
# =============================================================================
|
||||
|
||||
class User(Base):
|
||||
"""
|
||||
User model synced from Authentik
|
||||
|
||||
Users are created/updated when they authenticate via OIDC.
|
||||
The authentik_id links to the Authentik user record.
|
||||
"""
|
||||
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
authentik_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
email: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
nullable=False,
|
||||
)
|
||||
avatar_url: Mapped[str | None] = mapped_column(
|
||||
String(500),
|
||||
nullable=True,
|
||||
)
|
||||
api_keys_enabled: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=True,
|
||||
nullable=False,
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
last_login: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
roles: Mapped[List["Role"]] = relationship(
|
||||
"Role",
|
||||
secondary="user_roles",
|
||||
back_populates="users",
|
||||
lazy="selectin",
|
||||
)
|
||||
preferences: Mapped["UserPreferences"] = relationship(
|
||||
"UserPreferences",
|
||||
back_populates="user",
|
||||
uselist=False,
|
||||
lazy="selectin",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
api_keys: Mapped[List["ApiKey"]] = relationship(
|
||||
"ApiKey",
|
||||
back_populates="user",
|
||||
lazy="selectin",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<User {self.email}>"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Role Models
|
||||
# =============================================================================
|
||||
|
||||
class Role(Base):
|
||||
"""
|
||||
Role model for domain-scoped permissions
|
||||
|
||||
Permission format: domain.category:action
|
||||
- domain: Main area (control-room, library, media, ai, etc.)
|
||||
- category: Sub-area within domain (general for full access, or specific tools)
|
||||
- action: Permission level (viewer, user, editor, admin)
|
||||
|
||||
Roles are seeded from configuration, not user-editable.
|
||||
Groups are assigned roles via the group_roles mapping table.
|
||||
|
||||
Domains: control-room, library, media, ai, housekeeper, developer, documents, gaming, admin
|
||||
Actions: viewer, user, editor, admin (hierarchical)
|
||||
"""
|
||||
|
||||
__tablename__ = "roles"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(100),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Role name in format domain.category:action (e.g., control-room.general:admin)",
|
||||
)
|
||||
domain: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Permission domain (e.g., control-room, media, ai)",
|
||||
)
|
||||
category: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
nullable=False,
|
||||
default="general",
|
||||
comment="Permission category within domain (general for full access, or specific tool)",
|
||||
)
|
||||
action: Mapped[str] = mapped_column(
|
||||
String(20),
|
||||
nullable=False,
|
||||
comment="Permission action (viewer, user, editor, admin)",
|
||||
)
|
||||
|
||||
# Relationships
|
||||
users: Mapped[List["User"]] = relationship(
|
||||
"User",
|
||||
secondary="user_roles",
|
||||
back_populates="roles",
|
||||
lazy="selectin",
|
||||
)
|
||||
groups: Mapped[List["Group"]] = relationship(
|
||||
"Group",
|
||||
secondary="group_roles",
|
||||
back_populates="roles",
|
||||
lazy="selectin",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Role {self.name}>"
|
||||
|
||||
|
||||
class UserRole(Base):
|
||||
"""
|
||||
Association table for User-Role many-to-many relationship
|
||||
|
||||
Synced from Authentik groups during user authentication.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_roles"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
role_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("roles.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Group Model
|
||||
# =============================================================================
|
||||
|
||||
class Group(Base):
|
||||
"""
|
||||
Group model synced from Authentik
|
||||
|
||||
Groups are fetched from Authentik admin API and cached locally.
|
||||
They represent organizational units for access control.
|
||||
"""
|
||||
|
||||
__tablename__ = "groups"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
authentik_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Authentik group UUID",
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
unique=True,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
is_superuser: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=False,
|
||||
nullable=False,
|
||||
comment="Whether members of this group have superuser privileges",
|
||||
)
|
||||
parent_name: Mapped[str | None] = mapped_column(
|
||||
String(255),
|
||||
nullable=True,
|
||||
comment="Parent group name for hierarchy",
|
||||
)
|
||||
member_count: Mapped[int] = mapped_column(
|
||||
default=0,
|
||||
nullable=False,
|
||||
comment="Number of users in this group",
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
synced_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
comment="Last sync from Authentik",
|
||||
)
|
||||
|
||||
# Relationships
|
||||
roles: Mapped[List["Role"]] = relationship(
|
||||
"Role",
|
||||
secondary="group_roles",
|
||||
back_populates="groups",
|
||||
lazy="selectin",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Group {self.name}>"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# User Preferences Model
|
||||
# =============================================================================
|
||||
|
||||
class UserPreferences(Base):
|
||||
"""
|
||||
User preferences model
|
||||
|
||||
Stores user-specific settings that persist across sessions.
|
||||
Extended settings stored in preferences_json for flexibility.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_preferences"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
theme: Mapped[str] = mapped_column(
|
||||
String(20),
|
||||
default="system",
|
||||
nullable=False,
|
||||
comment="Theme preference: system, light, dark",
|
||||
)
|
||||
default_room: Mapped[str] = mapped_column(
|
||||
String(50),
|
||||
default="front-hall",
|
||||
nullable=False,
|
||||
comment="Default room for housekeeping features",
|
||||
)
|
||||
preferences_json: Mapped[dict] = mapped_column(
|
||||
JSONB,
|
||||
default=dict,
|
||||
nullable=False,
|
||||
comment="Extended preferences as JSON",
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User",
|
||||
back_populates="preferences",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<UserPreferences user_id={self.user_id}>"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# API Key Model
|
||||
# =============================================================================
|
||||
|
||||
class ApiKey(Base):
|
||||
"""
|
||||
API Key model for programmatic access
|
||||
|
||||
API keys provide an alternative to OIDC for:
|
||||
- Local development without SSO
|
||||
- Service-to-service communication
|
||||
- Scripts and automation
|
||||
|
||||
Keys inherit the user's roles but can optionally
|
||||
be restricted to a subset of scopes.
|
||||
"""
|
||||
|
||||
__tablename__ = "api_keys"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(100),
|
||||
nullable=False,
|
||||
comment="Human-readable key name (e.g., 'Dev Laptop', 'CI/CD')",
|
||||
)
|
||||
key_hash: Mapped[str] = mapped_column(
|
||||
String(255),
|
||||
nullable=False,
|
||||
comment="SHA-256 hash of the API key",
|
||||
)
|
||||
key_prefix: Mapped[str] = mapped_column(
|
||||
String(8),
|
||||
nullable=False,
|
||||
comment="First 8 chars of key for identification (e.g., 'cak_abc1')",
|
||||
)
|
||||
scopes: Mapped[List[str] | None] = mapped_column(
|
||||
ARRAY(String),
|
||||
nullable=True,
|
||||
comment="Optional scope restriction (subset of user roles)",
|
||||
)
|
||||
expires_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
comment="Optional expiration timestamp",
|
||||
)
|
||||
last_used_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=True,
|
||||
comment="Last time this key was used",
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User",
|
||||
back_populates="api_keys",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<ApiKey {self.key_prefix}... ({self.name})>"
|
||||
|
||||
@property
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if the API key has expired"""
|
||||
if self.expires_at is None:
|
||||
return False
|
||||
return datetime.now(self.expires_at.tzinfo) > self.expires_at
|
||||
@@ -0,0 +1,688 @@
|
||||
"""
|
||||
OIDC Authentication Module
|
||||
|
||||
Provides OAuth2/OIDC token validation for FastAPI using Authentik as IdP.
|
||||
Implements bearer token authentication with JWT verification.
|
||||
|
||||
Permission Format: domain.category:action
|
||||
- domain: Main area (control-room, library, media, ai, etc.)
|
||||
- category: Sub-area within domain (general for full domain, or specific tools)
|
||||
- action: Permission level (viewer, user, editor, admin)
|
||||
|
||||
Examples:
|
||||
- control-room.general:admin - Full access to Control Room
|
||||
- media.general:viewer - View-only access to Media area
|
||||
- ai.ollama:user - User-level access to Ollama specifically (future)
|
||||
|
||||
Action Hierarchy (higher implies lower):
|
||||
- admin > editor > user > viewer
|
||||
"""
|
||||
from fastapi import Depends, HTTPException, Security, Request
|
||||
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
|
||||
from jose import jwt, JWTError
|
||||
import httpx
|
||||
from functools import lru_cache
|
||||
from typing import Callable, Dict, List, Optional
|
||||
from src.shared.logging import get_logger
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Permission System
|
||||
# =============================================================================
|
||||
|
||||
# Action hierarchy: higher actions imply lower ones
|
||||
ACTION_HIERARCHY: Dict[str, int] = {
|
||||
"viewer": 1,
|
||||
"user": 2,
|
||||
"editor": 3,
|
||||
"admin": 4,
|
||||
}
|
||||
|
||||
# Valid domains (main areas)
|
||||
VALID_DOMAINS = {
|
||||
"control-room",
|
||||
"library",
|
||||
"media",
|
||||
"ai",
|
||||
"housekeeper",
|
||||
"developer",
|
||||
"documents",
|
||||
"gaming",
|
||||
"admin", # Global admin domain
|
||||
}
|
||||
|
||||
# Default category for general domain access
|
||||
DEFAULT_CATEGORY = "general"
|
||||
|
||||
logger = get_logger(__name__)
|
||||
security = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
class OIDCConfig:
|
||||
"""OIDC configuration from environment"""
|
||||
|
||||
def __init__(self):
|
||||
# These will be set from environment variables in config.py
|
||||
self.enabled = False
|
||||
self.issuer = ""
|
||||
self.audience = ""
|
||||
self.jwks_uri = ""
|
||||
|
||||
def configure(self, enabled: bool, issuer: str, audience: str):
|
||||
"""Configure OIDC settings"""
|
||||
self.enabled = enabled
|
||||
self.issuer = issuer
|
||||
self.audience = audience
|
||||
self.jwks_uri = f"{issuer.rstrip('/')}/jwks/"
|
||||
logger.info(f"OIDC configured: enabled={enabled}, issuer={issuer}")
|
||||
|
||||
|
||||
# Global OIDC config instance
|
||||
oidc_config = OIDCConfig()
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_jwks() -> Dict:
|
||||
"""
|
||||
Fetch JSON Web Key Set (JWKS) from Authentik
|
||||
|
||||
Cached to avoid repeated requests. Cache is cleared on server restart.
|
||||
|
||||
Returns:
|
||||
JWKS dictionary containing public keys for token verification
|
||||
|
||||
Raises:
|
||||
HTTPException: If JWKS fetch fails
|
||||
"""
|
||||
if not oidc_config.enabled:
|
||||
return {}
|
||||
|
||||
try:
|
||||
logger.debug(f"Fetching JWKS from {oidc_config.jwks_uri}")
|
||||
response = httpx.get(oidc_config.jwks_uri, timeout=10.0)
|
||||
response.raise_for_status()
|
||||
jwks = response.json()
|
||||
logger.info(f"JWKS fetched successfully ({len(jwks.get('keys', []))} keys)")
|
||||
return jwks
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch JWKS: {e}")
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="Authentication service unavailable"
|
||||
)
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Security(security)
|
||||
) -> Optional[Dict]:
|
||||
"""
|
||||
Validate OIDC token from Authorization: Bearer header
|
||||
|
||||
Extracts and validates JWT token from request header. Verifies:
|
||||
- Token signature using JWKS
|
||||
- Token expiration
|
||||
- Issuer matches Authentik
|
||||
- Audience matches core-api
|
||||
|
||||
Args:
|
||||
credentials: HTTP Bearer token from Authorization header
|
||||
|
||||
Returns:
|
||||
User claims dictionary containing email, name, groups, etc.
|
||||
Returns None if OIDC is disabled (allows unauthenticated access)
|
||||
|
||||
Raises:
|
||||
HTTPException 401: If token is invalid, expired, or missing when OIDC enabled
|
||||
"""
|
||||
# If OIDC is disabled, return a default local user
|
||||
if not oidc_config.enabled:
|
||||
logger.debug("OIDC disabled - using local user")
|
||||
return {
|
||||
"sub": "local-user",
|
||||
"email": "local@localhost",
|
||||
"preferred_username": "local",
|
||||
"name": "Local User",
|
||||
"groups": ["admin"],
|
||||
"auth_method": "local"
|
||||
}
|
||||
|
||||
# OIDC enabled - token required
|
||||
if not credentials:
|
||||
logger.warning("Authentication required but no token provided")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authentication required",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
token = credentials.credentials
|
||||
|
||||
try:
|
||||
# Decode token header to get key ID
|
||||
unverified_header = jwt.get_unverified_header(token)
|
||||
kid = unverified_header.get("kid")
|
||||
|
||||
if not kid:
|
||||
raise HTTPException(status_code=401, detail="Invalid token format")
|
||||
|
||||
# Find matching key in JWKS
|
||||
jwks = get_jwks()
|
||||
rsa_key = None
|
||||
|
||||
for key in jwks.get("keys", []):
|
||||
if key.get("kid") == kid:
|
||||
rsa_key = key
|
||||
break
|
||||
|
||||
if not rsa_key:
|
||||
logger.warning(f"No matching key found for kid: {kid}")
|
||||
raise HTTPException(status_code=401, detail="Invalid token key")
|
||||
|
||||
# Verify and decode token
|
||||
payload = jwt.decode(
|
||||
token,
|
||||
rsa_key,
|
||||
algorithms=["RS256"],
|
||||
audience=oidc_config.audience,
|
||||
issuer=oidc_config.issuer,
|
||||
)
|
||||
|
||||
user_email = payload.get("email", "unknown")
|
||||
logger.info(f"Authenticated user: {user_email}")
|
||||
|
||||
return payload
|
||||
|
||||
except jwt.ExpiredSignatureError:
|
||||
logger.warning("Token expired")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Token expired",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
except jwt.JWTClaimsError as e:
|
||||
logger.warning(f"Invalid token claims: {e}")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid token claims",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
except JWTError as e:
|
||||
logger.error(f"JWT validation error: {e}")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid authentication token",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected authentication error: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="Authentication error",
|
||||
)
|
||||
|
||||
|
||||
async def get_admin_user(
|
||||
user: Optional[Dict] = Depends(get_current_user)
|
||||
) -> Dict:
|
||||
"""
|
||||
Require admin group membership
|
||||
|
||||
Use this dependency for endpoints that require admin access.
|
||||
Checks if user is member of 'admin' group in Authentik.
|
||||
|
||||
Args:
|
||||
user: User claims from get_current_user
|
||||
|
||||
Returns:
|
||||
User claims dictionary if user is admin
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If user is not in admin group
|
||||
HTTPException 401: If OIDC enabled but user not authenticated
|
||||
"""
|
||||
# If OIDC disabled, allow all (backward compatibility)
|
||||
if not oidc_config.enabled or user is None:
|
||||
logger.debug("OIDC disabled - allowing admin access")
|
||||
return {"email": "unauthenticated", "groups": ["admin"]}
|
||||
|
||||
# Check admin group membership
|
||||
groups = user.get("groups", [])
|
||||
|
||||
if "admin" not in groups and "authentik Admins" not in groups:
|
||||
user_email = user.get("email", "unknown")
|
||||
logger.warning(f"User {user_email} attempted admin access (groups: {groups})")
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required"
|
||||
)
|
||||
|
||||
return user
|
||||
|
||||
|
||||
async def get_optional_user(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Security(security)
|
||||
) -> Optional[Dict]:
|
||||
"""
|
||||
Optional authentication - allows both authenticated and unauthenticated access
|
||||
|
||||
Use for endpoints that should be accessible to everyone but can provide
|
||||
enhanced functionality for authenticated users.
|
||||
|
||||
Args:
|
||||
credentials: HTTP Bearer token from Authorization header
|
||||
|
||||
Returns:
|
||||
User claims if valid token provided, local user if OIDC disabled, None otherwise
|
||||
"""
|
||||
# If OIDC is disabled, return the local user
|
||||
if not oidc_config.enabled:
|
||||
return {
|
||||
"sub": "local-user",
|
||||
"email": "local@localhost",
|
||||
"preferred_username": "local",
|
||||
"name": "Local User",
|
||||
"groups": ["admin"],
|
||||
"auth_method": "local"
|
||||
}
|
||||
|
||||
if not credentials:
|
||||
return None
|
||||
|
||||
try:
|
||||
return await get_current_user(credentials)
|
||||
except HTTPException:
|
||||
# Invalid token - return None instead of raising
|
||||
return None
|
||||
|
||||
|
||||
async def get_forward_auth_user(
|
||||
request: Request
|
||||
) -> Optional[Dict]:
|
||||
"""
|
||||
Authentik Forward Auth authentication for external access via NPM
|
||||
|
||||
This dependency allows:
|
||||
- External access through api.schweitz.net (with Authentik forward auth headers) - REQUIRES authentication
|
||||
- Internal direct access (no forward auth headers) - ALLOWED without authentication
|
||||
|
||||
When accessing through NPM with Authentik forward auth enabled, NPM adds headers like:
|
||||
- X-authentik-username
|
||||
- X-authentik-email
|
||||
- X-authentik-groups
|
||||
- X-authentik-name
|
||||
- X-authentik-uid
|
||||
|
||||
Args:
|
||||
request: FastAPI request object containing headers
|
||||
|
||||
Returns:
|
||||
User info dict if authenticated via forward auth headers
|
||||
None if accessed internally (no forward auth headers)
|
||||
|
||||
Raises:
|
||||
HTTPException 401: If forward auth headers present but invalid/incomplete
|
||||
"""
|
||||
# Check for Authentik forward auth headers
|
||||
username = request.headers.get("x-authentik-username")
|
||||
email = request.headers.get("x-authentik-email")
|
||||
groups = request.headers.get("x-authentik-groups")
|
||||
name = request.headers.get("x-authentik-name")
|
||||
uid = request.headers.get("x-authentik-uid")
|
||||
|
||||
# If NO forward auth headers present, this is internal access - allow it
|
||||
if not username and not email:
|
||||
logger.debug("No forward auth headers - allowing internal access")
|
||||
return None
|
||||
|
||||
# Forward auth headers present (external access via api.schweitz.net)
|
||||
# Validate authentication
|
||||
if not username or not email:
|
||||
logger.warning("Incomplete forward auth headers detected")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authentication required - incomplete forward auth headers"
|
||||
)
|
||||
|
||||
# Parse groups (comma-separated string to list)
|
||||
groups_list = [g.strip() for g in groups.split(",")] if groups else []
|
||||
|
||||
user_info = {
|
||||
"username": username,
|
||||
"email": email,
|
||||
"name": name or username,
|
||||
"groups": groups_list,
|
||||
"uid": uid,
|
||||
"auth_method": "forward_auth"
|
||||
}
|
||||
|
||||
logger.info(f"Authenticated via forward auth: {email} (groups: {groups_list})")
|
||||
return user_info
|
||||
|
||||
|
||||
async def get_forward_auth_admin(
|
||||
user: Optional[Dict] = Depends(get_forward_auth_user)
|
||||
) -> Dict:
|
||||
"""
|
||||
Require admin access for external requests, allow all internal requests
|
||||
|
||||
Use this dependency for endpoints that require admin access when accessed
|
||||
externally through api.schweitz.net, but allow unrestricted internal access.
|
||||
|
||||
Args:
|
||||
user: User info from get_forward_auth_user
|
||||
|
||||
Returns:
|
||||
User info dict if user is admin or if accessed internally
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If external user is not in admin/authentik Admins group
|
||||
"""
|
||||
# Internal access (no forward auth headers) - allow all
|
||||
if user is None:
|
||||
logger.debug("Internal access - allowing without admin check")
|
||||
return {"email": "internal", "groups": ["admin"], "auth_method": "internal"}
|
||||
|
||||
# External access - check admin group membership
|
||||
groups = user.get("groups", [])
|
||||
|
||||
if "admin" not in groups and "authentik Admins" not in groups:
|
||||
user_email = user.get("email", "unknown")
|
||||
logger.warning(f"User {user_email} attempted admin access (groups: {groups})")
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required"
|
||||
)
|
||||
|
||||
return user
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Permission-Based Access Control
|
||||
# =============================================================================
|
||||
|
||||
def _parse_permission(permission: str) -> tuple[str, str, str]:
|
||||
"""
|
||||
Parse a permission string into (domain, category, action)
|
||||
|
||||
Supports formats:
|
||||
- domain.category:action (full): "control-room.general:admin"
|
||||
- domain:action (shorthand): "control-room:admin" -> ("control-room", "general", "admin")
|
||||
|
||||
Returns:
|
||||
Tuple of (domain, category, action)
|
||||
|
||||
Raises:
|
||||
ValueError: If permission format is invalid
|
||||
"""
|
||||
# Split on colon first to get action
|
||||
if ":" not in permission:
|
||||
raise ValueError(f"Invalid permission format (missing ':'): {permission}")
|
||||
|
||||
location, action = permission.rsplit(":", 1)
|
||||
|
||||
# Split location on dot to get domain and category
|
||||
if "." in location:
|
||||
domain, category = location.split(".", 1)
|
||||
else:
|
||||
# Shorthand: domain:action -> domain.general:action
|
||||
domain = location
|
||||
category = DEFAULT_CATEGORY
|
||||
|
||||
return domain, category, action
|
||||
|
||||
|
||||
def _action_satisfies(user_action: str, required_action: str) -> bool:
|
||||
"""
|
||||
Check if user's action level satisfies the required action
|
||||
|
||||
Due to hierarchy, admin satisfies editor, editor satisfies user, etc.
|
||||
|
||||
Args:
|
||||
user_action: The action the user has
|
||||
required_action: The action required for access
|
||||
|
||||
Returns:
|
||||
True if user's action is >= required action
|
||||
"""
|
||||
user_level = ACTION_HIERARCHY.get(user_action, 0)
|
||||
required_level = ACTION_HIERARCHY.get(required_action, 0)
|
||||
return user_level >= required_level
|
||||
|
||||
|
||||
def _user_has_permission(
|
||||
user_permissions: List[str],
|
||||
required_domain: str,
|
||||
required_category: str,
|
||||
required_action: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if user has a permission that satisfies the requirement
|
||||
|
||||
Checks:
|
||||
1. Exact match: domain.category:action
|
||||
2. Domain-wide: domain.general:action (if category != general)
|
||||
3. Global admin: admin.general:admin (superuser)
|
||||
|
||||
Args:
|
||||
user_permissions: List of user's permission strings
|
||||
required_domain: Required domain
|
||||
required_category: Required category
|
||||
required_action: Required action
|
||||
|
||||
Returns:
|
||||
True if user has sufficient permission
|
||||
"""
|
||||
for perm in user_permissions:
|
||||
try:
|
||||
dom, cat, act = _parse_permission(perm)
|
||||
except ValueError:
|
||||
continue
|
||||
|
||||
# Global admin (admin.general:admin) grants all permissions
|
||||
if dom == "admin" and cat == "general" and act == "admin":
|
||||
return True
|
||||
|
||||
# Check if this permission covers the requirement
|
||||
if dom == required_domain:
|
||||
# Exact category match
|
||||
if cat == required_category and _action_satisfies(act, required_action):
|
||||
return True
|
||||
# Domain-wide permission (general category) covers all categories in domain
|
||||
if cat == "general" and _action_satisfies(act, required_action):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _extract_permissions_from_groups(groups: List[str]) -> List[str]:
|
||||
"""
|
||||
Extract permission strings from Authentik group names
|
||||
|
||||
Authentik groups follow naming: tatlock-{domain}-{category}-{action}
|
||||
or shorthand: tatlock-{domain}-{action} (implies category=general)
|
||||
|
||||
Examples:
|
||||
- tatlock-control-room-general-admin -> control-room.general:admin
|
||||
- tatlock-media-viewer -> media.general:viewer (shorthand)
|
||||
- tatlock-ai-ollama-user -> ai.ollama:user
|
||||
|
||||
Args:
|
||||
groups: List of Authentik group names
|
||||
|
||||
Returns:
|
||||
List of permission strings
|
||||
"""
|
||||
permissions = []
|
||||
|
||||
for group in groups:
|
||||
if not group.startswith("tatlock-"):
|
||||
continue
|
||||
|
||||
# Remove prefix
|
||||
parts = group[8:].split("-") # Remove "tatlock-"
|
||||
|
||||
if len(parts) >= 3:
|
||||
# Could be domain-category-action or domain-with-hyphen-action
|
||||
# Try to find a valid action at the end
|
||||
action = parts[-1]
|
||||
if action in ACTION_HIERARCHY:
|
||||
# Check if domain-category or single domain with hyphen
|
||||
remaining = parts[:-1]
|
||||
|
||||
# Try to find known domain (greedy match from start)
|
||||
for i in range(len(remaining), 0, -1):
|
||||
potential_domain = "-".join(remaining[:i])
|
||||
if potential_domain in VALID_DOMAINS:
|
||||
category_parts = remaining[i:]
|
||||
category = "-".join(category_parts) if category_parts else DEFAULT_CATEGORY
|
||||
permissions.append(f"{potential_domain}.{category}:{action}")
|
||||
break
|
||||
elif len(parts) == 2:
|
||||
# Shorthand: domain-action (domain might have hyphen)
|
||||
action = parts[-1]
|
||||
if action in ACTION_HIERARCHY:
|
||||
domain = parts[0]
|
||||
if domain in VALID_DOMAINS:
|
||||
permissions.append(f"{domain}.{DEFAULT_CATEGORY}:{action}")
|
||||
|
||||
return permissions
|
||||
|
||||
|
||||
def require_permission(
|
||||
domain: str,
|
||||
action: str,
|
||||
category: str = DEFAULT_CATEGORY,
|
||||
) -> Callable:
|
||||
"""
|
||||
Dependency factory for permission-based access control
|
||||
|
||||
Creates a FastAPI dependency that checks if the current user has
|
||||
the required permission. Considers action hierarchy and global admin.
|
||||
|
||||
Usage:
|
||||
@router.get("/containers")
|
||||
async def list_containers(
|
||||
user: Dict = Depends(require_permission("control-room", "viewer"))
|
||||
):
|
||||
...
|
||||
|
||||
@router.delete("/container/{id}")
|
||||
async def delete_container(
|
||||
user: Dict = Depends(require_permission("control-room", "admin"))
|
||||
):
|
||||
...
|
||||
|
||||
Args:
|
||||
domain: Permission domain (e.g., "control-room", "media")
|
||||
action: Required action level (viewer, user, editor, admin)
|
||||
category: Permission category within domain, defaults to "general"
|
||||
|
||||
Returns:
|
||||
FastAPI dependency function
|
||||
"""
|
||||
perm_str = f"{domain}.{category}:{action}"
|
||||
|
||||
async def permission_checker(
|
||||
user: Optional[Dict] = Depends(get_current_user)
|
||||
) -> Dict:
|
||||
"""Check if user has required permission"""
|
||||
|
||||
# If OIDC disabled, allow all (local dev mode)
|
||||
if not oidc_config.enabled:
|
||||
logger.debug(f"OIDC disabled - allowing {perm_str}")
|
||||
return user or {"email": "local", "groups": ["admin"]}
|
||||
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authentication required",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
# Extract permissions from user's groups
|
||||
groups = user.get("groups", [])
|
||||
permissions = _extract_permissions_from_groups(groups)
|
||||
|
||||
# Check if user has required permission
|
||||
if _user_has_permission(permissions, domain, category, action):
|
||||
logger.debug(f"User {user.get('email')} granted {perm_str}")
|
||||
return user
|
||||
|
||||
# Permission denied
|
||||
user_email = user.get("email", "unknown")
|
||||
logger.warning(
|
||||
f"User {user_email} denied {perm_str} "
|
||||
f"(groups: {groups}, permissions: {permissions})"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"Permission required: {perm_str}",
|
||||
)
|
||||
|
||||
return permission_checker
|
||||
|
||||
|
||||
def require_any_permission(*required_permissions: str) -> Callable:
|
||||
"""
|
||||
Dependency factory requiring any one of multiple permissions
|
||||
|
||||
Useful for endpoints accessible to multiple roles.
|
||||
|
||||
Usage:
|
||||
@router.get("/shared-resource")
|
||||
async def get_shared(
|
||||
user: Dict = Depends(require_any_permission(
|
||||
"control-room:viewer",
|
||||
"media:viewer",
|
||||
))
|
||||
):
|
||||
...
|
||||
|
||||
Args:
|
||||
*required_permissions: Permission strings (domain.category:action or domain:action)
|
||||
|
||||
Returns:
|
||||
FastAPI dependency function
|
||||
"""
|
||||
|
||||
async def permission_checker(
|
||||
user: Optional[Dict] = Depends(get_current_user)
|
||||
) -> Dict:
|
||||
"""Check if user has any of the required permissions"""
|
||||
|
||||
# If OIDC disabled, allow all
|
||||
if not oidc_config.enabled:
|
||||
return user or {"email": "local", "groups": ["admin"]}
|
||||
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authentication required",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
groups = user.get("groups", [])
|
||||
permissions = _extract_permissions_from_groups(groups)
|
||||
|
||||
# Check each required permission
|
||||
for perm in required_permissions:
|
||||
try:
|
||||
dom, cat, act = _parse_permission(perm)
|
||||
if _user_has_permission(permissions, dom, cat, act):
|
||||
logger.debug(f"User {user.get('email')} granted via {perm}")
|
||||
return user
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid permission format: {perm}")
|
||||
continue
|
||||
|
||||
# None matched
|
||||
user_email = user.get("email", "unknown")
|
||||
logger.warning(
|
||||
f"User {user_email} denied (required any of: {required_permissions})"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"One of these permissions required: {', '.join(required_permissions)}",
|
||||
)
|
||||
|
||||
return permission_checker
|
||||
@@ -0,0 +1,211 @@
|
||||
"""
|
||||
Authentication Schemas
|
||||
|
||||
Pydantic models for auth request/response payloads.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from pydantic import Field
|
||||
|
||||
from src.shared.base import BaseSchema
|
||||
|
||||
|
||||
class AuthSyncRequest(BaseSchema):
|
||||
"""
|
||||
Request payload for POST /auth/sync
|
||||
|
||||
The client sends this after obtaining an OIDC token from Authentik.
|
||||
The access_token is validated against Authentik's userinfo endpoint.
|
||||
"""
|
||||
|
||||
access_token: str = Field(
|
||||
...,
|
||||
description="OIDC access token from Authentik",
|
||||
)
|
||||
|
||||
|
||||
class RoleSchema(BaseSchema):
|
||||
"""Role information in domain.category:action format"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Role ID")
|
||||
name: str = Field(..., description="Role name (e.g., 'control-room.general:admin')")
|
||||
domain: str = Field(..., description="Permission domain (e.g., 'control-room')")
|
||||
category: str = Field(default="general", description="Permission category (e.g., 'general')")
|
||||
action: str = Field(..., description="Permission action (e.g., 'admin')")
|
||||
|
||||
|
||||
class UserPreferencesSchema(BaseSchema):
|
||||
"""User preferences"""
|
||||
|
||||
theme: str = Field(default="system", description="Theme preference: system, light, dark")
|
||||
default_room: str = Field(default="front-hall", description="Default room for housekeeping")
|
||||
preferences_json: dict = Field(default_factory=dict, description="Extended preferences")
|
||||
|
||||
|
||||
class UserSchema(BaseSchema):
|
||||
"""User information returned from sync"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal user ID")
|
||||
authentik_id: uuid.UUID = Field(..., description="Authentik user ID")
|
||||
email: str = Field(..., description="User email")
|
||||
name: str = Field(..., description="Display name")
|
||||
avatar_url: Optional[str] = Field(None, description="Profile picture URL")
|
||||
created_at: datetime = Field(..., description="Account creation timestamp")
|
||||
last_login: Optional[datetime] = Field(None, description="Last login timestamp")
|
||||
|
||||
|
||||
class AuthSyncResponse(BaseSchema):
|
||||
"""
|
||||
Response from POST /auth/sync
|
||||
|
||||
Contains the synced user profile, roles, and preferences.
|
||||
"""
|
||||
|
||||
user: UserSchema = Field(..., description="User profile")
|
||||
roles: list[RoleSchema] = Field(..., description="User's permission roles")
|
||||
preferences: UserPreferencesSchema = Field(..., description="User preferences")
|
||||
is_new_user: bool = Field(..., description="True if user was just created")
|
||||
|
||||
|
||||
class TokenInfoSchema(BaseSchema):
|
||||
"""
|
||||
Token information from Authentik userinfo endpoint
|
||||
|
||||
This is what Authentik returns when validating an access token.
|
||||
"""
|
||||
|
||||
sub: str = Field(..., description="Subject (Authentik user ID)")
|
||||
email: str = Field(..., description="User email")
|
||||
name: Optional[str] = Field(None, description="Display name")
|
||||
preferred_username: Optional[str] = Field(None, description="Username")
|
||||
groups: list[str] = Field(default_factory=list, description="Group memberships")
|
||||
picture: Optional[str] = Field(None, description="Profile picture URL")
|
||||
|
||||
|
||||
class UserListItemSchema(BaseSchema):
|
||||
"""User item for list display"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal user ID")
|
||||
email: str = Field(..., description="User email")
|
||||
name: str = Field(..., description="Display name")
|
||||
avatar_url: Optional[str] = Field(None, description="Profile picture URL")
|
||||
created_at: datetime = Field(..., description="Account creation timestamp")
|
||||
last_login: Optional[datetime] = Field(None, description="Last login timestamp")
|
||||
roles: list[str] = Field(default_factory=list, description="Role names")
|
||||
|
||||
|
||||
class UsersListResponse(BaseSchema):
|
||||
"""Response from GET /auth/users"""
|
||||
|
||||
items: list[UserListItemSchema] = Field(..., description="List of users")
|
||||
total: int = Field(..., description="Total count of users")
|
||||
|
||||
|
||||
class BulkSyncResultSchema(BaseSchema):
|
||||
"""Result from bulk sync operation"""
|
||||
|
||||
created: int = Field(..., description="Number of users created")
|
||||
updated: int = Field(..., description="Number of users updated")
|
||||
failed: int = Field(..., description="Number of users that failed to sync")
|
||||
total_in_authentik: int = Field(..., description="Total users in Authentik")
|
||||
errors: list[str] = Field(default_factory=list, description="Error messages for failed syncs")
|
||||
|
||||
|
||||
class GroupListItemSchema(BaseSchema):
|
||||
"""Group item for list display"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="Internal group ID")
|
||||
authentik_id: uuid.UUID = Field(..., description="Authentik group ID")
|
||||
name: str = Field(..., description="Group name")
|
||||
is_superuser: bool = Field(default=False, description="Whether group has superuser privileges")
|
||||
parent_name: Optional[str] = Field(None, description="Parent group name")
|
||||
member_count: int = Field(default=0, description="Number of users in this group")
|
||||
synced_at: datetime = Field(..., description="Last sync timestamp")
|
||||
roles: list[str] = Field(default_factory=list, description="Assigned role names")
|
||||
|
||||
|
||||
class GroupsListResponse(BaseSchema):
|
||||
"""Response from GET /auth/groups"""
|
||||
|
||||
items: list[GroupListItemSchema] = Field(..., description="List of groups")
|
||||
total: int = Field(..., description="Total count of groups")
|
||||
|
||||
|
||||
class RolesListResponse(BaseSchema):
|
||||
"""Response from GET /auth/roles"""
|
||||
|
||||
items: list[RoleSchema] = Field(..., description="List of all roles")
|
||||
total: int = Field(..., description="Total count of roles")
|
||||
|
||||
|
||||
class GroupRoleAssignmentResponse(BaseSchema):
|
||||
"""Response from group role assignment operations"""
|
||||
|
||||
group_id: uuid.UUID = Field(..., description="Group ID")
|
||||
group_name: str = Field(..., description="Group name")
|
||||
roles: list[str] = Field(..., description="Currently assigned role names")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# User Profile (Phase 4)
|
||||
# =============================================================================
|
||||
|
||||
class UserProfileResponse(BaseSchema):
|
||||
"""Response from GET /users/me - full user profile"""
|
||||
|
||||
user: UserSchema = Field(..., description="User profile")
|
||||
roles: list[RoleSchema] = Field(..., description="User's permission roles")
|
||||
preferences: UserPreferencesSchema = Field(..., description="User preferences")
|
||||
|
||||
|
||||
class PreferencesUpdateRequest(BaseSchema):
|
||||
"""Request for PATCH /users/me/preferences"""
|
||||
|
||||
theme: Optional[str] = Field(None, description="Theme preference: system, light, dark")
|
||||
default_room: Optional[str] = Field(None, description="Default room for housekeeping")
|
||||
preferences_json: Optional[dict] = Field(None, description="Extended preferences (merged)")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# API Keys (Phase 4)
|
||||
# =============================================================================
|
||||
|
||||
class ApiKeyCreateRequest(BaseSchema):
|
||||
"""Request for POST /users/me/api-keys"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=100, description="Human-readable key name")
|
||||
scopes: Optional[list[str]] = Field(None, description="Optional scope restriction (role names)")
|
||||
expires_in_days: Optional[int] = Field(None, ge=1, le=365, description="Days until expiration (optional)")
|
||||
|
||||
|
||||
class ApiKeyCreateResponse(BaseSchema):
|
||||
"""Response from POST /users/me/api-keys - includes the key (shown only once)"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="API key ID")
|
||||
name: str = Field(..., description="Key name")
|
||||
key: str = Field(..., description="The API key (shown only once!)")
|
||||
key_prefix: str = Field(..., description="Key prefix for identification")
|
||||
scopes: Optional[list[str]] = Field(None, description="Scope restriction")
|
||||
expires_at: Optional[datetime] = Field(None, description="Expiration timestamp")
|
||||
created_at: datetime = Field(..., description="Creation timestamp")
|
||||
|
||||
|
||||
class ApiKeySchema(BaseSchema):
|
||||
"""API key information (without the actual key)"""
|
||||
|
||||
id: uuid.UUID = Field(..., description="API key ID")
|
||||
name: str = Field(..., description="Key name")
|
||||
key_prefix: str = Field(..., description="Key prefix for identification (e.g., 'tak_abc1')")
|
||||
scopes: Optional[list[str]] = Field(None, description="Scope restriction")
|
||||
expires_at: Optional[datetime] = Field(None, description="Expiration timestamp")
|
||||
last_used_at: Optional[datetime] = Field(None, description="Last usage timestamp")
|
||||
created_at: datetime = Field(..., description="Creation timestamp")
|
||||
is_expired: bool = Field(..., description="Whether the key has expired")
|
||||
|
||||
|
||||
class ApiKeysListResponse(BaseSchema):
|
||||
"""Response from GET /users/me/api-keys"""
|
||||
|
||||
items: list[ApiKeySchema] = Field(..., description="List of API keys")
|
||||
total: int = Field(..., description="Total count of keys")
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Dashboard Domain
|
||||
|
||||
Provides dashboard management endpoints including quick links.
|
||||
"""
|
||||
from src.domains.dashboard.controller import dashboard_controller
|
||||
|
||||
__all__ = ["dashboard_controller"]
|
||||
@@ -0,0 +1,305 @@
|
||||
"""
|
||||
Dashboard Controller
|
||||
|
||||
Provides API endpoints for dashboard management including quick links.
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, Depends, Query
|
||||
from typing import Dict, Optional
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.database import get_async_session
|
||||
from src.shared.logging import get_logger
|
||||
from src.domains.auth.oidc import get_current_user, get_optional_user
|
||||
from src.domains.dashboard.service import get_dashboard_service
|
||||
from src.domains.dashboard.schemas import (
|
||||
QuickLinkCreate,
|
||||
QuickLinkUpdate,
|
||||
QuickLinkResponse,
|
||||
QuickLinkListResponse,
|
||||
QuickLinkReorderRequest,
|
||||
QuickLinkReorderResponse,
|
||||
DashboardWidgetCreate,
|
||||
DashboardWidgetUpdate,
|
||||
DashboardWidgetResponse,
|
||||
DashboardWidgetListResponse,
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class DashboardController(BaseController):
|
||||
"""
|
||||
Controller for dashboard operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Quick links CRUD
|
||||
- Quick links reordering
|
||||
- Dashboard widgets management
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/dashboard", tags=["Dashboard"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
service = get_dashboard_service()
|
||||
|
||||
# =====================================================================
|
||||
# Quick Links
|
||||
# =====================================================================
|
||||
|
||||
@router.get(
|
||||
"/quick-links",
|
||||
response_model=QuickLinkListResponse,
|
||||
summary="List quick links"
|
||||
)
|
||||
async def list_quick_links(
|
||||
category: Optional[str] = Query(None, description="Filter by category"),
|
||||
include_global: bool = Query(True, description="Include global links"),
|
||||
visible_only: bool = Query(True, description="Only visible links"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Optional[Dict] = Depends(get_optional_user),
|
||||
):
|
||||
"""
|
||||
List quick links for the current user
|
||||
|
||||
Returns user-specific links plus global links (if include_global=True).
|
||||
"""
|
||||
user_id = user.get("sub") if user else None
|
||||
|
||||
links = await service.get_quick_links(
|
||||
session=session,
|
||||
user_id=user_id,
|
||||
include_global=include_global,
|
||||
category=category,
|
||||
visible_only=visible_only,
|
||||
)
|
||||
|
||||
return QuickLinkListResponse(
|
||||
links=[QuickLinkResponse.model_validate(link, from_attributes=True) for link in links],
|
||||
total=len(links)
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/quick-links/{link_id}",
|
||||
response_model=QuickLinkResponse,
|
||||
summary="Get a quick link"
|
||||
)
|
||||
async def get_quick_link(
|
||||
link_id: int,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Optional[Dict] = Depends(get_optional_user),
|
||||
):
|
||||
"""Get a specific quick link by ID"""
|
||||
user_id = user.get("sub") if user else None
|
||||
|
||||
link = await service.get_quick_link(session, link_id, user_id)
|
||||
if not link:
|
||||
raise HTTPException(status_code=404, detail="Quick link not found")
|
||||
|
||||
return QuickLinkResponse.model_validate(link, from_attributes=True)
|
||||
|
||||
@router.post(
|
||||
"/quick-links",
|
||||
response_model=QuickLinkResponse,
|
||||
status_code=201,
|
||||
summary="Create a quick link"
|
||||
)
|
||||
async def create_quick_link(
|
||||
data: QuickLinkCreate,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""
|
||||
Create a new quick link for the current user
|
||||
|
||||
Links are user-specific by default. Admins can create global links
|
||||
by setting user_id to null.
|
||||
"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
link = await service.create_quick_link(session, data, user_id)
|
||||
|
||||
logger.info(f"Quick link created: {link.title} by user {user.get('preferred_username')}")
|
||||
|
||||
return QuickLinkResponse.model_validate(link, from_attributes=True)
|
||||
|
||||
@router.put(
|
||||
"/quick-links/{link_id}",
|
||||
response_model=QuickLinkResponse,
|
||||
summary="Update a quick link"
|
||||
)
|
||||
async def update_quick_link(
|
||||
link_id: int,
|
||||
data: QuickLinkUpdate,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""Update an existing quick link"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
link = await service.update_quick_link(session, link_id, data, user_id)
|
||||
if not link:
|
||||
raise HTTPException(status_code=404, detail="Quick link not found or not authorized")
|
||||
|
||||
return QuickLinkResponse.model_validate(link, from_attributes=True)
|
||||
|
||||
@router.delete(
|
||||
"/quick-links/{link_id}",
|
||||
status_code=204,
|
||||
summary="Delete a quick link"
|
||||
)
|
||||
async def delete_quick_link(
|
||||
link_id: int,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""Delete a quick link"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
success = await service.delete_quick_link(session, link_id, user_id)
|
||||
if not success:
|
||||
raise HTTPException(status_code=404, detail="Quick link not found or not authorized")
|
||||
|
||||
return None
|
||||
|
||||
@router.post(
|
||||
"/quick-links/reorder",
|
||||
response_model=QuickLinkReorderResponse,
|
||||
summary="Reorder quick links"
|
||||
)
|
||||
async def reorder_quick_links(
|
||||
data: QuickLinkReorderRequest,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""
|
||||
Reorder quick links by providing link IDs in desired order
|
||||
|
||||
The position of each link will be set to its index in the provided list.
|
||||
"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
reordered = await service.reorder_quick_links(session, data.link_ids, user_id)
|
||||
|
||||
return QuickLinkReorderResponse(
|
||||
success=True,
|
||||
message=f"Reordered {reordered} links",
|
||||
reordered_count=reordered
|
||||
)
|
||||
|
||||
# =====================================================================
|
||||
# Dashboard Widgets
|
||||
# =====================================================================
|
||||
|
||||
@router.get(
|
||||
"/widgets",
|
||||
response_model=DashboardWidgetListResponse,
|
||||
summary="List dashboard widgets"
|
||||
)
|
||||
async def list_widgets(
|
||||
include_defaults: bool = Query(True, description="Include default widgets"),
|
||||
visible_only: bool = Query(True, description="Only visible widgets"),
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Optional[Dict] = Depends(get_optional_user),
|
||||
):
|
||||
"""List dashboard widgets for the current user"""
|
||||
user_id = user.get("sub") if user else None
|
||||
|
||||
widgets = await service.get_widgets(
|
||||
session=session,
|
||||
user_id=user_id,
|
||||
include_defaults=include_defaults,
|
||||
visible_only=visible_only,
|
||||
)
|
||||
|
||||
return DashboardWidgetListResponse(
|
||||
widgets=[DashboardWidgetResponse.model_validate(w, from_attributes=True) for w in widgets],
|
||||
total=len(widgets)
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/widgets/{widget_id}",
|
||||
response_model=DashboardWidgetResponse,
|
||||
summary="Get a dashboard widget"
|
||||
)
|
||||
async def get_widget(
|
||||
widget_id: int,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Optional[Dict] = Depends(get_optional_user),
|
||||
):
|
||||
"""Get a specific dashboard widget by ID"""
|
||||
user_id = user.get("sub") if user else None
|
||||
|
||||
widget = await service.get_widget(session, widget_id, user_id)
|
||||
if not widget:
|
||||
raise HTTPException(status_code=404, detail="Widget not found")
|
||||
|
||||
return DashboardWidgetResponse.model_validate(widget, from_attributes=True)
|
||||
|
||||
@router.post(
|
||||
"/widgets",
|
||||
response_model=DashboardWidgetResponse,
|
||||
status_code=201,
|
||||
summary="Create a dashboard widget"
|
||||
)
|
||||
async def create_widget(
|
||||
data: DashboardWidgetCreate,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""Create a new dashboard widget"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
widget = await service.create_widget(session, data, user_id)
|
||||
|
||||
logger.info(f"Widget created: {widget.widget_type} by user {user.get('preferred_username')}")
|
||||
|
||||
return DashboardWidgetResponse.model_validate(widget, from_attributes=True)
|
||||
|
||||
@router.put(
|
||||
"/widgets/{widget_id}",
|
||||
response_model=DashboardWidgetResponse,
|
||||
summary="Update a dashboard widget"
|
||||
)
|
||||
async def update_widget(
|
||||
widget_id: int,
|
||||
data: DashboardWidgetUpdate,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""Update an existing dashboard widget"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
widget = await service.update_widget(session, widget_id, data, user_id)
|
||||
if not widget:
|
||||
raise HTTPException(status_code=404, detail="Widget not found or not authorized")
|
||||
|
||||
return DashboardWidgetResponse.model_validate(widget, from_attributes=True)
|
||||
|
||||
@router.delete(
|
||||
"/widgets/{widget_id}",
|
||||
status_code=204,
|
||||
summary="Delete a dashboard widget"
|
||||
)
|
||||
async def delete_widget(
|
||||
widget_id: int,
|
||||
session: AsyncSession = Depends(get_async_session),
|
||||
user: Dict = Depends(get_current_user),
|
||||
):
|
||||
"""Delete a dashboard widget"""
|
||||
user_id = user.get("sub")
|
||||
|
||||
success = await service.delete_widget(session, widget_id, user_id)
|
||||
if not success:
|
||||
raise HTTPException(status_code=404, detail="Widget not found or not authorized")
|
||||
|
||||
return None
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
dashboard_controller = DashboardController()
|
||||
@@ -0,0 +1,77 @@
|
||||
"""
|
||||
Dashboard Domain Models
|
||||
|
||||
SQLAlchemy models for dashboard-related data.
|
||||
"""
|
||||
from datetime import datetime
|
||||
from sqlalchemy import Column, Integer, String, Boolean, DateTime, Text, ForeignKey
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from src.shared.database import Base
|
||||
|
||||
|
||||
class QuickLink(Base):
|
||||
"""Quick link for dashboard jump pad"""
|
||||
__tablename__ = "quick_links"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Link content
|
||||
title = Column(String(100), nullable=False)
|
||||
url = Column(String(500), nullable=False)
|
||||
icon = Column(String(100), nullable=True) # Icon name or URL
|
||||
description = Column(String(255), nullable=True)
|
||||
|
||||
# Categorization
|
||||
category = Column(String(50), nullable=True) # e.g., "services", "tools", "docs"
|
||||
|
||||
# User association - nullable for global links
|
||||
user_id = Column(String(255), nullable=True, index=True) # Authentik user ID
|
||||
|
||||
# Ordering and display
|
||||
position = Column(Integer, default=0)
|
||||
is_visible = Column(Boolean, default=True)
|
||||
link_type = Column(String(20), default="iframe") # "iframe", "new_tab", etc.
|
||||
|
||||
# Styling
|
||||
color = Column(String(20), nullable=True) # Hex color for the link card
|
||||
background_color = Column(String(20), nullable=True)
|
||||
|
||||
# Metadata
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<QuickLink(id={self.id}, title='{self.title}', user_id='{self.user_id}')>"
|
||||
|
||||
|
||||
class DashboardWidget(Base):
|
||||
"""Dashboard widget configuration"""
|
||||
__tablename__ = "dashboard_widgets"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Widget identification
|
||||
widget_type = Column(String(50), nullable=False) # e.g., "quick_links", "service_status", "weather"
|
||||
|
||||
# User association - nullable for default widgets
|
||||
user_id = Column(String(255), nullable=True, index=True)
|
||||
|
||||
# Position and sizing
|
||||
position_x = Column(Integer, default=0)
|
||||
position_y = Column(Integer, default=0)
|
||||
width = Column(Integer, default=1)
|
||||
height = Column(Integer, default=1)
|
||||
|
||||
# Widget-specific configuration (JSON)
|
||||
config = Column(Text, nullable=True) # JSON string for widget-specific settings
|
||||
|
||||
# Display
|
||||
is_visible = Column(Boolean, default=True)
|
||||
|
||||
# Metadata
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<DashboardWidget(id={self.id}, type='{self.widget_type}', user_id='{self.user_id}')>"
|
||||
@@ -0,0 +1,113 @@
|
||||
"""
|
||||
Dashboard Domain Schemas
|
||||
|
||||
Pydantic schemas for dashboard endpoints.
|
||||
"""
|
||||
from datetime import datetime
|
||||
from typing import Optional, List
|
||||
from pydantic import Field
|
||||
|
||||
from src.shared.base import BaseSchema
|
||||
|
||||
|
||||
# Quick Link Schemas
|
||||
class QuickLinkBase(BaseSchema):
|
||||
"""Base schema for quick links"""
|
||||
title: str = Field(..., min_length=1, max_length=100, description="Link title")
|
||||
url: str = Field(..., min_length=1, max_length=500, description="Link URL")
|
||||
icon: Optional[str] = Field(None, max_length=100, description="Icon name or URL")
|
||||
description: Optional[str] = Field(None, max_length=255, description="Link description")
|
||||
category: Optional[str] = Field(None, max_length=50, description="Link category")
|
||||
color: Optional[str] = Field(None, max_length=20, description="Hex color for link card")
|
||||
background_color: Optional[str] = Field(None, max_length=20, description="Background hex color")
|
||||
|
||||
|
||||
class QuickLinkCreate(QuickLinkBase):
|
||||
"""Schema for creating a quick link"""
|
||||
position: Optional[int] = Field(0, ge=0, description="Display position")
|
||||
is_visible: Optional[bool] = Field(True, description="Whether link is visible")
|
||||
link_type: Optional[str] = Field("iframe", max_length=20, description="Link type: iframe, new_tab")
|
||||
|
||||
|
||||
class QuickLinkUpdate(BaseSchema):
|
||||
"""Schema for updating a quick link"""
|
||||
title: Optional[str] = Field(None, min_length=1, max_length=100)
|
||||
url: Optional[str] = Field(None, min_length=1, max_length=500)
|
||||
icon: Optional[str] = Field(None, max_length=100)
|
||||
description: Optional[str] = Field(None, max_length=255)
|
||||
category: Optional[str] = Field(None, max_length=50)
|
||||
position: Optional[int] = Field(None, ge=0)
|
||||
is_visible: Optional[bool] = None
|
||||
link_type: Optional[str] = Field(None, max_length=20)
|
||||
color: Optional[str] = Field(None, max_length=20)
|
||||
background_color: Optional[str] = Field(None, max_length=20)
|
||||
|
||||
|
||||
class QuickLinkResponse(QuickLinkBase):
|
||||
"""Schema for quick link response"""
|
||||
id: int
|
||||
user_id: Optional[str] = None
|
||||
position: int
|
||||
is_visible: bool
|
||||
link_type: str = "iframe"
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class QuickLinkListResponse(BaseSchema):
|
||||
"""Response for list of quick links"""
|
||||
links: List[QuickLinkResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class QuickLinkReorderRequest(BaseSchema):
|
||||
"""Request to reorder quick links"""
|
||||
link_ids: List[int] = Field(..., description="List of link IDs in desired order")
|
||||
|
||||
|
||||
class QuickLinkReorderResponse(BaseSchema):
|
||||
"""Response after reordering"""
|
||||
success: bool
|
||||
message: str
|
||||
reordered_count: int
|
||||
|
||||
|
||||
# Dashboard Widget Schemas
|
||||
class DashboardWidgetBase(BaseSchema):
|
||||
"""Base schema for dashboard widgets"""
|
||||
widget_type: str = Field(..., min_length=1, max_length=50, description="Widget type identifier")
|
||||
position_x: int = Field(0, ge=0, description="X position on grid")
|
||||
position_y: int = Field(0, ge=0, description="Y position on grid")
|
||||
width: int = Field(1, ge=1, le=12, description="Widget width in grid units")
|
||||
height: int = Field(1, ge=1, le=12, description="Widget height in grid units")
|
||||
config: Optional[str] = Field(None, description="JSON config for widget")
|
||||
is_visible: bool = Field(True, description="Whether widget is visible")
|
||||
|
||||
|
||||
class DashboardWidgetCreate(DashboardWidgetBase):
|
||||
"""Schema for creating a widget"""
|
||||
pass
|
||||
|
||||
|
||||
class DashboardWidgetUpdate(BaseSchema):
|
||||
"""Schema for updating a widget"""
|
||||
position_x: Optional[int] = Field(None, ge=0)
|
||||
position_y: Optional[int] = Field(None, ge=0)
|
||||
width: Optional[int] = Field(None, ge=1, le=12)
|
||||
height: Optional[int] = Field(None, ge=1, le=12)
|
||||
config: Optional[str] = None
|
||||
is_visible: Optional[bool] = None
|
||||
|
||||
|
||||
class DashboardWidgetResponse(DashboardWidgetBase):
|
||||
"""Schema for widget response"""
|
||||
id: int
|
||||
user_id: Optional[str] = None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class DashboardWidgetListResponse(BaseSchema):
|
||||
"""Response for list of widgets"""
|
||||
widgets: List[DashboardWidgetResponse]
|
||||
total: int
|
||||
@@ -0,0 +1,319 @@
|
||||
"""
|
||||
Dashboard Domain Service
|
||||
|
||||
Business logic for dashboard operations.
|
||||
"""
|
||||
from typing import Optional, List
|
||||
from sqlalchemy import select, update, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.shared.logging import get_logger
|
||||
from src.domains.dashboard.models import QuickLink, DashboardWidget
|
||||
from src.domains.dashboard.schemas import (
|
||||
QuickLinkCreate,
|
||||
QuickLinkUpdate,
|
||||
QuickLinkResponse,
|
||||
DashboardWidgetCreate,
|
||||
DashboardWidgetUpdate,
|
||||
DashboardWidgetResponse,
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class DashboardService:
|
||||
"""Service for dashboard operations"""
|
||||
|
||||
# =========================================================================
|
||||
# Quick Links
|
||||
# =========================================================================
|
||||
|
||||
async def get_quick_links(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
user_id: Optional[str] = None,
|
||||
include_global: bool = True,
|
||||
category: Optional[str] = None,
|
||||
visible_only: bool = True,
|
||||
) -> List[QuickLink]:
|
||||
"""
|
||||
Get quick links for a user
|
||||
|
||||
Args:
|
||||
session: Database session
|
||||
user_id: User ID to filter by (None for global only)
|
||||
include_global: Whether to include global links (user_id=None)
|
||||
category: Optional category filter
|
||||
visible_only: Only return visible links
|
||||
"""
|
||||
conditions = []
|
||||
|
||||
if user_id:
|
||||
if include_global:
|
||||
from sqlalchemy import or_
|
||||
conditions.append(or_(QuickLink.user_id == user_id, QuickLink.user_id.is_(None)))
|
||||
else:
|
||||
conditions.append(QuickLink.user_id == user_id)
|
||||
else:
|
||||
conditions.append(QuickLink.user_id.is_(None))
|
||||
|
||||
if category:
|
||||
conditions.append(QuickLink.category == category)
|
||||
|
||||
if visible_only:
|
||||
conditions.append(QuickLink.is_visible == True)
|
||||
|
||||
stmt = select(QuickLink).where(*conditions).order_by(QuickLink.position, QuickLink.id)
|
||||
result = await session.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def get_quick_link(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
link_id: int,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Optional[QuickLink]:
|
||||
"""Get a specific quick link by ID"""
|
||||
conditions = [QuickLink.id == link_id]
|
||||
|
||||
if user_id:
|
||||
from sqlalchemy import or_
|
||||
conditions.append(or_(QuickLink.user_id == user_id, QuickLink.user_id.is_(None)))
|
||||
|
||||
stmt = select(QuickLink).where(*conditions)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def create_quick_link(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
data: QuickLinkCreate,
|
||||
user_id: Optional[str] = None,
|
||||
) -> QuickLink:
|
||||
"""Create a new quick link"""
|
||||
# Get max position for this user
|
||||
stmt = select(QuickLink.position).where(
|
||||
QuickLink.user_id == user_id if user_id else QuickLink.user_id.is_(None)
|
||||
).order_by(QuickLink.position.desc()).limit(1)
|
||||
result = await session.execute(stmt)
|
||||
max_pos = result.scalar_one_or_none() or -1
|
||||
|
||||
link = QuickLink(
|
||||
title=data.title,
|
||||
url=data.url,
|
||||
icon=data.icon,
|
||||
description=data.description,
|
||||
category=data.category,
|
||||
position=data.position if data.position > 0 else max_pos + 1,
|
||||
is_visible=data.is_visible,
|
||||
color=data.color,
|
||||
background_color=data.background_color,
|
||||
user_id=user_id,
|
||||
)
|
||||
session.add(link)
|
||||
await session.commit()
|
||||
await session.refresh(link)
|
||||
|
||||
logger.info(f"Created quick link: {link.title} (id={link.id}, user={user_id})")
|
||||
return link
|
||||
|
||||
async def update_quick_link(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
link_id: int,
|
||||
data: QuickLinkUpdate,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Optional[QuickLink]:
|
||||
"""Update a quick link"""
|
||||
link = await self.get_quick_link(session, link_id, user_id)
|
||||
if not link:
|
||||
return None
|
||||
|
||||
# Only allow updating own links or global links for admins
|
||||
if link.user_id and link.user_id != user_id:
|
||||
return None
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for field, value in update_data.items():
|
||||
setattr(link, field, value)
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(link)
|
||||
|
||||
logger.info(f"Updated quick link: {link.title} (id={link.id})")
|
||||
return link
|
||||
|
||||
async def delete_quick_link(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
link_id: int,
|
||||
user_id: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Delete a quick link"""
|
||||
link = await self.get_quick_link(session, link_id, user_id)
|
||||
if not link:
|
||||
return False
|
||||
|
||||
# Only allow deleting own links
|
||||
if link.user_id and link.user_id != user_id:
|
||||
return False
|
||||
|
||||
await session.delete(link)
|
||||
await session.commit()
|
||||
|
||||
logger.info(f"Deleted quick link: id={link_id}")
|
||||
return True
|
||||
|
||||
async def reorder_quick_links(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
link_ids: List[int],
|
||||
user_id: Optional[str] = None,
|
||||
) -> int:
|
||||
"""Reorder quick links by updating positions"""
|
||||
reordered = 0
|
||||
|
||||
for position, link_id in enumerate(link_ids):
|
||||
conditions = [QuickLink.id == link_id]
|
||||
if user_id:
|
||||
from sqlalchemy import or_
|
||||
conditions.append(or_(QuickLink.user_id == user_id, QuickLink.user_id.is_(None)))
|
||||
|
||||
stmt = update(QuickLink).where(*conditions).values(position=position)
|
||||
result = await session.execute(stmt)
|
||||
reordered += result.rowcount
|
||||
|
||||
await session.commit()
|
||||
|
||||
logger.info(f"Reordered {reordered} quick links for user {user_id}")
|
||||
return reordered
|
||||
|
||||
# =========================================================================
|
||||
# Dashboard Widgets
|
||||
# =========================================================================
|
||||
|
||||
async def get_widgets(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
user_id: Optional[str] = None,
|
||||
include_defaults: bool = True,
|
||||
visible_only: bool = True,
|
||||
) -> List[DashboardWidget]:
|
||||
"""Get dashboard widgets for a user"""
|
||||
conditions = []
|
||||
|
||||
if user_id:
|
||||
if include_defaults:
|
||||
from sqlalchemy import or_
|
||||
conditions.append(or_(DashboardWidget.user_id == user_id, DashboardWidget.user_id.is_(None)))
|
||||
else:
|
||||
conditions.append(DashboardWidget.user_id == user_id)
|
||||
else:
|
||||
conditions.append(DashboardWidget.user_id.is_(None))
|
||||
|
||||
if visible_only:
|
||||
conditions.append(DashboardWidget.is_visible == True)
|
||||
|
||||
stmt = select(DashboardWidget).where(*conditions).order_by(
|
||||
DashboardWidget.position_y, DashboardWidget.position_x
|
||||
)
|
||||
result = await session.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def get_widget(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
widget_id: int,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Optional[DashboardWidget]:
|
||||
"""Get a specific widget by ID"""
|
||||
conditions = [DashboardWidget.id == widget_id]
|
||||
|
||||
if user_id:
|
||||
from sqlalchemy import or_
|
||||
conditions.append(or_(DashboardWidget.user_id == user_id, DashboardWidget.user_id.is_(None)))
|
||||
|
||||
stmt = select(DashboardWidget).where(*conditions)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def create_widget(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
data: DashboardWidgetCreate,
|
||||
user_id: Optional[str] = None,
|
||||
) -> DashboardWidget:
|
||||
"""Create a new dashboard widget"""
|
||||
widget = DashboardWidget(
|
||||
widget_type=data.widget_type,
|
||||
position_x=data.position_x,
|
||||
position_y=data.position_y,
|
||||
width=data.width,
|
||||
height=data.height,
|
||||
config=data.config,
|
||||
is_visible=data.is_visible,
|
||||
user_id=user_id,
|
||||
)
|
||||
session.add(widget)
|
||||
await session.commit()
|
||||
await session.refresh(widget)
|
||||
|
||||
logger.info(f"Created widget: {widget.widget_type} (id={widget.id}, user={user_id})")
|
||||
return widget
|
||||
|
||||
async def update_widget(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
widget_id: int,
|
||||
data: DashboardWidgetUpdate,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Optional[DashboardWidget]:
|
||||
"""Update a dashboard widget"""
|
||||
widget = await self.get_widget(session, widget_id, user_id)
|
||||
if not widget:
|
||||
return None
|
||||
|
||||
if widget.user_id and widget.user_id != user_id:
|
||||
return None
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for field, value in update_data.items():
|
||||
setattr(widget, field, value)
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(widget)
|
||||
|
||||
logger.info(f"Updated widget: id={widget.id}")
|
||||
return widget
|
||||
|
||||
async def delete_widget(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
widget_id: int,
|
||||
user_id: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Delete a dashboard widget"""
|
||||
widget = await self.get_widget(session, widget_id, user_id)
|
||||
if not widget:
|
||||
return False
|
||||
|
||||
if widget.user_id and widget.user_id != user_id:
|
||||
return False
|
||||
|
||||
await session.delete(widget)
|
||||
await session.commit()
|
||||
|
||||
logger.info(f"Deleted widget: id={widget_id}")
|
||||
return True
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_dashboard_service: Optional[DashboardService] = None
|
||||
|
||||
|
||||
def get_dashboard_service() -> DashboardService:
|
||||
"""Get singleton dashboard service instance"""
|
||||
global _dashboard_service
|
||||
if _dashboard_service is None:
|
||||
_dashboard_service = DashboardService()
|
||||
return _dashboard_service
|
||||
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Health Domain
|
||||
|
||||
Provides health check and diagnostics endpoints.
|
||||
"""
|
||||
from src.domains.health.controller import health_controller
|
||||
|
||||
__all__ = ["health_controller"]
|
||||
@@ -0,0 +1,228 @@
|
||||
"""
|
||||
Health Controller
|
||||
|
||||
Provides service health and information endpoints
|
||||
"""
|
||||
from fastapi import APIRouter, Response
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.config import get_settings
|
||||
from src.shared.logging import get_logger
|
||||
from src.shared.database import get_database
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class HealthController(BaseController):
|
||||
"""
|
||||
Controller for service health and information
|
||||
|
||||
Provides endpoints for:
|
||||
- Service information and status
|
||||
- Health checks
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="", tags=["Health"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(tags=self.tags)
|
||||
settings = get_settings()
|
||||
|
||||
@router.get(
|
||||
"/",
|
||||
summary="Service information",
|
||||
response_class=JSONResponse
|
||||
)
|
||||
async def root():
|
||||
"""
|
||||
Get service information and health status
|
||||
|
||||
Returns basic information about the API service and available endpoints.
|
||||
"""
|
||||
logger.debug("Root endpoint accessed")
|
||||
return {
|
||||
"service": settings.app_name,
|
||||
"version": settings.app_version,
|
||||
"status": "healthy",
|
||||
"docs": "/docs"
|
||||
}
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
summary="Health check",
|
||||
response_class=JSONResponse
|
||||
)
|
||||
async def health_check():
|
||||
"""
|
||||
Fast health check endpoint for container orchestration
|
||||
|
||||
Returns a 200 OK immediately if the service is running.
|
||||
Does NOT check backend connectivity (use /health/full for that).
|
||||
Used by Docker, Kubernetes, and load balancers for liveness probes.
|
||||
"""
|
||||
return {
|
||||
"status": "healthy",
|
||||
"version": settings.app_version
|
||||
}
|
||||
|
||||
@router.get(
|
||||
"/health/full",
|
||||
summary="Fast health check for Docker",
|
||||
)
|
||||
async def full_health_check(response: Response):
|
||||
"""
|
||||
Fast health check for container orchestration (Docker/K8s).
|
||||
|
||||
Checks component availability WITHOUT running expensive operations.
|
||||
Returns 200 OK if all components are available, otherwise 503.
|
||||
|
||||
For detailed diagnostics, use /health/diagnostics instead.
|
||||
"""
|
||||
import time
|
||||
start_time = time.time()
|
||||
|
||||
# Import here to avoid circular imports
|
||||
from src.models.ollama_client import get_ollama_client
|
||||
|
||||
# Check 1: Ollama connection + verify agent model is available
|
||||
ollama_client = get_ollama_client()
|
||||
ollama_healthy = False
|
||||
ollama_error = None
|
||||
model_available = False
|
||||
|
||||
try:
|
||||
# Ping Ollama
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
|
||||
# Verify the agent model is pulled and check what's currently loaded
|
||||
models_info = {}
|
||||
if ollama_healthy:
|
||||
try:
|
||||
models_response = await ollama_client.list_models()
|
||||
available_models = [m.get('name', '') for m in models_response.get('models', [])]
|
||||
model_available = settings.agent_model in available_models
|
||||
|
||||
# Get info about currently loaded models (those with size in memory)
|
||||
loaded_models = [
|
||||
m.get('name', '') for m in models_response.get('models', [])
|
||||
if m.get('size', 0) > 0
|
||||
]
|
||||
|
||||
models_info = {
|
||||
"configured": settings.agent_model,
|
||||
"available": model_available,
|
||||
"total_in_ollama": len(available_models),
|
||||
"currently_loaded": loaded_models if loaded_models else ["none"]
|
||||
}
|
||||
|
||||
if not model_available:
|
||||
ollama_error = f"Model '{settings.agent_model}' not found in Ollama. Available: {', '.join(available_models[:3])}"
|
||||
ollama_healthy = False
|
||||
except Exception as e:
|
||||
ollama_error = f"Could not list Ollama models: {str(e)}"
|
||||
ollama_healthy = False
|
||||
|
||||
except Exception as e:
|
||||
ollama_error = str(e)
|
||||
logger.warning(f"Ollama health check failed: {ollama_error}")
|
||||
|
||||
# Check 2: Database connection
|
||||
database = get_database()
|
||||
db_healthy = False
|
||||
db_error = None
|
||||
|
||||
try:
|
||||
db_healthy = await database.health_check()
|
||||
except Exception as e:
|
||||
db_error = str(e)
|
||||
logger.warning(f"Database health check failed: {db_error}")
|
||||
|
||||
is_healthy = ollama_healthy and db_healthy
|
||||
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
status_code = 200 if is_healthy else 503
|
||||
response.status_code = status_code
|
||||
|
||||
return {
|
||||
"status": "healthy" if is_healthy else "unhealthy",
|
||||
"status_code": status_code,
|
||||
"response_time_ms": elapsed_ms,
|
||||
"components": {
|
||||
"ollama": {
|
||||
"status": "healthy" if ollama_healthy else "unhealthy",
|
||||
"models": models_info if models_info else {
|
||||
"configured": settings.agent_model,
|
||||
"available": False
|
||||
},
|
||||
"error": ollama_error
|
||||
},
|
||||
"database": {
|
||||
"status": "healthy" if db_healthy else "unhealthy",
|
||||
"error": db_error
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@router.get(
|
||||
"/health/diagnostics",
|
||||
summary="Detailed system diagnostics",
|
||||
)
|
||||
async def diagnostics(deep_test: bool = False):
|
||||
"""
|
||||
Comprehensive system diagnostics with detailed component information.
|
||||
|
||||
Query Parameters:
|
||||
- deep_test: Set to true to actually test agent generation (slow, ~5-10s)
|
||||
|
||||
Returns detailed information about all system components.
|
||||
"""
|
||||
import time
|
||||
from src.models.ollama_client import get_ollama_client
|
||||
|
||||
start_time = time.time()
|
||||
diagnostics = {
|
||||
"timestamp": time.time(),
|
||||
"service": {
|
||||
"name": settings.app_name,
|
||||
"version": settings.app_version,
|
||||
"purpose": "Infrastructure management and tools API"
|
||||
},
|
||||
"components": {}
|
||||
}
|
||||
|
||||
# 1. Ollama Connection
|
||||
ollama_client = get_ollama_client()
|
||||
try:
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
diagnostics["components"]["ollama"] = {
|
||||
"status": "connected",
|
||||
"url": settings.ollama_base_url,
|
||||
"timeout": settings.ollama_timeout,
|
||||
"default_model": settings.default_model
|
||||
}
|
||||
except Exception as e:
|
||||
diagnostics["components"]["ollama"] = {
|
||||
"status": "error",
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
# 2. Configuration
|
||||
diagnostics["configuration"] = {
|
||||
"agent_fallback_enabled": settings.agent_fallback_enabled,
|
||||
"memory_tier1_max_turns": settings.memory_tier1_max_turns,
|
||||
"cors_origins": settings.cors_origins[:2] if len(settings.cors_origins) > 2 else settings.cors_origins
|
||||
}
|
||||
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
diagnostics["response_time_ms"] = elapsed_ms
|
||||
|
||||
return diagnostics
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
health_controller = HealthController()
|
||||
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Housekeeping Domain
|
||||
|
||||
Provides home automation endpoints via Home Assistant.
|
||||
"""
|
||||
from src.domains.housekeeping.controller import housekeeping_controller
|
||||
|
||||
__all__ = ["housekeeping_controller"]
|
||||
@@ -0,0 +1,645 @@
|
||||
"""
|
||||
Housekeeping Controller
|
||||
|
||||
Provides API endpoints for home automation via Home Assistant.
|
||||
Designed for the Tatlock Housekeeper agent and other consumers.
|
||||
"""
|
||||
import asyncio
|
||||
from fastapi import APIRouter, HTTPException, Query, Depends
|
||||
from typing import List, Dict, Any, Optional
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.clients import get_homeassistant_client
|
||||
from src.shared.logging import get_logger
|
||||
from src.domains.auth.oidc import get_admin_user
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
# Pydantic Schemas
|
||||
class Device(BaseModel):
|
||||
"""Device/entity information"""
|
||||
entity_id: str
|
||||
name: str
|
||||
domain: str
|
||||
area: Optional[str] = None
|
||||
state: str
|
||||
attributes: Dict[str, Any] = {}
|
||||
last_changed: Optional[str] = None
|
||||
|
||||
|
||||
class DeviceListResponse(BaseModel):
|
||||
"""Response for device listing"""
|
||||
devices: List[Device]
|
||||
|
||||
|
||||
class DeviceDetailResponse(Device):
|
||||
"""Detailed device response"""
|
||||
pass
|
||||
|
||||
|
||||
class Area(BaseModel):
|
||||
"""Area/room information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class AreaListResponse(BaseModel):
|
||||
"""Response for area listing"""
|
||||
areas: List[Area]
|
||||
|
||||
|
||||
class DeviceControlRequest(BaseModel):
|
||||
"""Request to control a device"""
|
||||
action: str = Field(..., description="Action: turn_on, turn_off, or toggle")
|
||||
brightness: Optional[int] = Field(None, ge=0, le=255)
|
||||
color_temp: Optional[int] = None
|
||||
rgb_color: Optional[List[int]] = None
|
||||
|
||||
class Config:
|
||||
extra = "allow"
|
||||
|
||||
|
||||
class DeviceControlResponse(BaseModel):
|
||||
"""Response from device control"""
|
||||
success: bool
|
||||
entity_id: str
|
||||
new_state: Optional[str] = None
|
||||
message: str
|
||||
|
||||
|
||||
class Scene(BaseModel):
|
||||
"""Scene information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class SceneListResponse(BaseModel):
|
||||
"""Response for scene listing"""
|
||||
scenes: List[Scene]
|
||||
|
||||
|
||||
class SceneActivateResponse(BaseModel):
|
||||
"""Response from scene activation"""
|
||||
success: bool
|
||||
scene_id: str
|
||||
message: str
|
||||
|
||||
|
||||
class Script(BaseModel):
|
||||
"""Script information"""
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
class ScriptListResponse(BaseModel):
|
||||
"""Response for script listing"""
|
||||
scripts: List[Script]
|
||||
|
||||
|
||||
class ScriptRunRequest(BaseModel):
|
||||
"""Request to run a script"""
|
||||
variables: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ScriptRunResponse(BaseModel):
|
||||
"""Response from script execution"""
|
||||
success: bool
|
||||
script_id: str
|
||||
message: str
|
||||
|
||||
|
||||
class Automation(BaseModel):
|
||||
"""Automation information"""
|
||||
id: str
|
||||
name: str
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutomationListResponse(BaseModel):
|
||||
"""Response for automation listing"""
|
||||
automations: List[Automation]
|
||||
|
||||
|
||||
class AutomationToggleRequest(BaseModel):
|
||||
"""Request to toggle automation"""
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutomationToggleResponse(BaseModel):
|
||||
"""Response from automation toggle"""
|
||||
success: bool
|
||||
automation_id: str
|
||||
enabled: bool
|
||||
message: str
|
||||
|
||||
|
||||
class HistoryEntry(BaseModel):
|
||||
"""Single history entry"""
|
||||
state: str
|
||||
timestamp: str
|
||||
attributes: Dict[str, Any] = {}
|
||||
|
||||
|
||||
class HistoryResponse(BaseModel):
|
||||
"""Response for history query"""
|
||||
entity_id: str
|
||||
history: List[HistoryEntry]
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""Health check response"""
|
||||
status: str
|
||||
connected: bool
|
||||
platform: str
|
||||
version: Optional[str] = None
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
class ErrorResponse(BaseModel):
|
||||
"""Standard error response"""
|
||||
error: bool = True
|
||||
code: str
|
||||
message: str
|
||||
|
||||
|
||||
class HousekeepingController(BaseController):
|
||||
"""
|
||||
Controller for home automation operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Device discovery and control
|
||||
- Scene activation
|
||||
- Script execution
|
||||
- Automation management
|
||||
- State history
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/housekeeping", tags=["Housekeeping"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
response_model=HealthResponse,
|
||||
summary="Home automation health check"
|
||||
)
|
||||
async def get_health():
|
||||
"""Check Home Assistant connection health"""
|
||||
ha = get_homeassistant_client()
|
||||
return await ha.health_check()
|
||||
|
||||
@router.get(
|
||||
"/devices",
|
||||
response_model=DeviceListResponse,
|
||||
summary="List available devices"
|
||||
)
|
||||
async def list_devices(
|
||||
domain: Optional[str] = Query(None, description="Filter by domain (light, switch, climate, etc.)"),
|
||||
area: Optional[str] = Query(None, description="Filter by area/room name")
|
||||
):
|
||||
"""List all available devices with optional filtering"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
excluded_domains = {
|
||||
"zone", "person", "device_tracker", "sun", "weather",
|
||||
"persistent_notification", "update", "binary_sensor", "sensor",
|
||||
"conversation", "calendar", "button", "number", "select",
|
||||
"text", "time", "date", "datetime", "image", "tts", "stt"
|
||||
}
|
||||
|
||||
devices = []
|
||||
for state in states:
|
||||
entity_id = state.get("entity_id", "")
|
||||
entity_domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||
|
||||
if entity_domain in excluded_domains:
|
||||
continue
|
||||
|
||||
if domain and entity_domain != domain:
|
||||
continue
|
||||
|
||||
device_area = state.get("attributes", {}).get("area_id")
|
||||
|
||||
if area and device_area and area.lower() not in device_area.lower():
|
||||
continue
|
||||
|
||||
device = Device(
|
||||
entity_id=entity_id,
|
||||
name=state.get("attributes", {}).get("friendly_name", entity_id),
|
||||
domain=entity_domain,
|
||||
area=device_area,
|
||||
state=state.get("state", "unknown"),
|
||||
attributes=state.get("attributes", {}),
|
||||
last_changed=state.get("last_changed")
|
||||
)
|
||||
devices.append(device)
|
||||
|
||||
return DeviceListResponse(devices=devices)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list devices: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/devices/{entity_id:path}",
|
||||
response_model=DeviceDetailResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def get_device(entity_id: str):
|
||||
"""Get detailed state of a specific device"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
state = await ha.get_state(entity_id)
|
||||
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "DEVICE_NOT_FOUND",
|
||||
"message": f"Device {entity_id} not found"}
|
||||
)
|
||||
|
||||
entity_domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||
|
||||
return DeviceDetailResponse(
|
||||
entity_id=entity_id,
|
||||
name=state.get("attributes", {}).get("friendly_name", entity_id),
|
||||
domain=entity_domain,
|
||||
area=state.get("attributes", {}).get("area_id"),
|
||||
state=state.get("state", "unknown"),
|
||||
attributes=state.get("attributes", {}),
|
||||
last_changed=state.get("last_changed")
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get device {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/areas",
|
||||
response_model=AreaListResponse,
|
||||
summary="List areas/rooms"
|
||||
)
|
||||
async def list_areas():
|
||||
"""List all configured areas/rooms in Home Assistant"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
areas = await ha.get_areas()
|
||||
return AreaListResponse(
|
||||
areas=[Area(id=a["id"], name=a["name"]) for a in areas]
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list areas: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/devices/{entity_id:path}/control",
|
||||
response_model=DeviceControlResponse,
|
||||
responses={404: {"model": ErrorResponse}, 400: {"model": ErrorResponse}}
|
||||
)
|
||||
async def control_device(
|
||||
entity_id: str,
|
||||
request: DeviceControlRequest,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""Control a device (turn on, turn off, toggle, or set attributes)"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
valid_actions = ["turn_on", "turn_off", "toggle"]
|
||||
if request.action not in valid_actions:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": True, "code": "INVALID_ACTION",
|
||||
"message": f"Invalid action '{request.action}'. Must be one of: {', '.join(valid_actions)}"}
|
||||
)
|
||||
|
||||
try:
|
||||
current_state = await ha.get_state(entity_id)
|
||||
if not current_state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "DEVICE_NOT_FOUND",
|
||||
"message": f"Device {entity_id} not found"}
|
||||
)
|
||||
|
||||
attributes = {}
|
||||
if request.brightness is not None:
|
||||
attributes["brightness"] = request.brightness
|
||||
if request.color_temp is not None:
|
||||
attributes["color_temp"] = request.color_temp
|
||||
if request.rgb_color is not None:
|
||||
attributes["rgb_color"] = request.rgb_color
|
||||
|
||||
extra_fields = request.model_dump(exclude={"action", "brightness", "color_temp", "rgb_color"})
|
||||
for key, value in extra_fields.items():
|
||||
if value is not None:
|
||||
attributes[key] = value
|
||||
|
||||
if request.action == "turn_on":
|
||||
await ha.turn_on(entity_id, **attributes)
|
||||
elif request.action == "turn_off":
|
||||
await ha.turn_off(entity_id)
|
||||
else:
|
||||
await ha.toggle(entity_id)
|
||||
|
||||
await asyncio.sleep(0.3)
|
||||
|
||||
new_state = await ha.get_state(entity_id)
|
||||
|
||||
logger.info(f"Device {entity_id} controlled: {request.action} by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return DeviceControlResponse(
|
||||
success=True,
|
||||
entity_id=entity_id,
|
||||
new_state=new_state.get("state") if new_state else None,
|
||||
message=f"Device {request.action} successful"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to control device {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/scenes",
|
||||
response_model=SceneListResponse,
|
||||
summary="List available scenes"
|
||||
)
|
||||
async def list_scenes():
|
||||
"""List all available scenes in Home Assistant"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
scenes = [
|
||||
Scene(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"])
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("scene.")
|
||||
]
|
||||
|
||||
return SceneListResponse(scenes=scenes)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list scenes: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/scenes/{scene_id:path}/activate",
|
||||
response_model=SceneActivateResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def activate_scene(
|
||||
scene_id: str,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""Activate a scene"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
state = await ha.get_state(scene_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "SCENE_NOT_FOUND",
|
||||
"message": f"Scene {scene_id} not found"}
|
||||
)
|
||||
|
||||
await ha.activate_scene(scene_id)
|
||||
|
||||
logger.info(f"Scene {scene_id} activated by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return SceneActivateResponse(
|
||||
success=True,
|
||||
scene_id=scene_id,
|
||||
message="Scene activated"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to activate scene {scene_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/scripts",
|
||||
response_model=ScriptListResponse,
|
||||
summary="List available scripts"
|
||||
)
|
||||
async def list_scripts():
|
||||
"""List all available scripts/sequences in Home Assistant"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
scripts = [
|
||||
Script(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"])
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("script.")
|
||||
]
|
||||
|
||||
return ScriptListResponse(scripts=scripts)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list scripts: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/scripts/{script_id:path}/run",
|
||||
response_model=ScriptRunResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def run_script(
|
||||
script_id: str,
|
||||
request: Optional[ScriptRunRequest] = None,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""Execute a script with optional variables"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
state = await ha.get_state(script_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "SCRIPT_NOT_FOUND",
|
||||
"message": f"Script {script_id} not found"}
|
||||
)
|
||||
|
||||
variables = request.variables if request else None
|
||||
await ha.run_script(script_id, variables)
|
||||
|
||||
logger.info(f"Script {script_id} executed by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return ScriptRunResponse(
|
||||
success=True,
|
||||
script_id=script_id,
|
||||
message="Script executed"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to run script {script_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/automations",
|
||||
response_model=AutomationListResponse,
|
||||
summary="List automations"
|
||||
)
|
||||
async def list_automations():
|
||||
"""List all automations with their enabled/disabled status"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
states = await ha.get_states()
|
||||
|
||||
automations = [
|
||||
Automation(
|
||||
id=s["entity_id"],
|
||||
name=s.get("attributes", {}).get("friendly_name", s["entity_id"]),
|
||||
enabled=s.get("state") == "on"
|
||||
)
|
||||
for s in states
|
||||
if s["entity_id"].startswith("automation.")
|
||||
]
|
||||
|
||||
return AutomationListResponse(automations=automations)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list automations: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.post(
|
||||
"/automations/{automation_id:path}/toggle",
|
||||
response_model=AutomationToggleResponse,
|
||||
responses={404: {"model": ErrorResponse}}
|
||||
)
|
||||
async def toggle_automation(
|
||||
automation_id: str,
|
||||
request: AutomationToggleRequest,
|
||||
user: Dict = Depends(get_admin_user)
|
||||
):
|
||||
"""Enable or disable an automation"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
state = await ha.get_state(automation_id)
|
||||
if not state:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": True, "code": "AUTOMATION_NOT_FOUND",
|
||||
"message": f"Automation {automation_id} not found"}
|
||||
)
|
||||
|
||||
if request.enabled:
|
||||
await ha.enable_automation(automation_id)
|
||||
else:
|
||||
await ha.disable_automation(automation_id)
|
||||
|
||||
logger.info(f"Automation {automation_id} {'enabled' if request.enabled else 'disabled'} by {user.get('preferred_username', 'unknown')}")
|
||||
|
||||
return AutomationToggleResponse(
|
||||
success=True,
|
||||
automation_id=automation_id,
|
||||
enabled=request.enabled,
|
||||
message=f"Automation {'enabled' if request.enabled else 'disabled'}"
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to toggle automation {automation_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/history",
|
||||
response_model=HistoryResponse,
|
||||
responses={400: {"model": ErrorResponse}}
|
||||
)
|
||||
async def get_history(
|
||||
entity_id: str = Query(..., description="Entity ID to get history for"),
|
||||
hours: int = Query(24, ge=1, le=168, description="Hours of history (1-168)")
|
||||
):
|
||||
"""Get state history for a device"""
|
||||
ha = get_homeassistant_client()
|
||||
|
||||
try:
|
||||
history_data = await ha.get_history(entity_id, hours)
|
||||
|
||||
history_entries = []
|
||||
if history_data and len(history_data) > 0:
|
||||
for entry in history_data[0]:
|
||||
history_entries.append(HistoryEntry(
|
||||
state=entry.get("state", "unknown"),
|
||||
timestamp=entry.get("last_changed", ""),
|
||||
attributes=entry.get("attributes", {})
|
||||
))
|
||||
|
||||
return HistoryResponse(
|
||||
entity_id=entity_id,
|
||||
history=history_entries
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get history for {entity_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": True, "code": "CONNECTION_ERROR", "message": str(e)}
|
||||
)
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
housekeeping_controller = HousekeepingController()
|
||||
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Infrastructure Domain
|
||||
|
||||
Provides infrastructure management endpoints for Docker/Portainer and NPM.
|
||||
"""
|
||||
from src.domains.infrastructure.controller import infrastructure_controller
|
||||
|
||||
__all__ = ["infrastructure_controller"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Static Domain
|
||||
|
||||
Serves static files for widgets and other frontend assets.
|
||||
"""
|
||||
from src.domains.static.controller import static_controller
|
||||
|
||||
__all__ = ["static_controller"]
|
||||
@@ -0,0 +1,116 @@
|
||||
"""
|
||||
Static Files Controller
|
||||
|
||||
Serves static files for widgets and other frontend assets.
|
||||
"""
|
||||
from fastapi import APIRouter
|
||||
from fastapi.responses import FileResponse, HTMLResponse
|
||||
from pathlib import Path
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class StaticController(BaseController):
|
||||
"""
|
||||
Controller for serving static files
|
||||
|
||||
Provides endpoints for:
|
||||
- Organizr widgets
|
||||
- Other static assets
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/static", tags=["Static"])
|
||||
# Static files are at the root of the project
|
||||
self.static_dir = Path(__file__).parent.parent.parent.parent / "static"
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.get(
|
||||
"/widgets/{filename}",
|
||||
response_class=HTMLResponse,
|
||||
summary="Get widget file"
|
||||
)
|
||||
async def get_widget(filename: str):
|
||||
"""
|
||||
Serve widget HTML files
|
||||
|
||||
Args:
|
||||
filename: Widget filename (e.g., service-control.html)
|
||||
|
||||
Returns:
|
||||
HTML file content
|
||||
"""
|
||||
widget_path = self.static_dir / "widgets" / filename
|
||||
|
||||
if not widget_path.exists():
|
||||
return HTMLResponse(
|
||||
content=f"<h1>404 - Widget not found</h1><p>{filename}</p>",
|
||||
status_code=404
|
||||
)
|
||||
|
||||
if not widget_path.is_file():
|
||||
return HTMLResponse(
|
||||
content=f"<h1>400 - Not a file</h1>",
|
||||
status_code=400
|
||||
)
|
||||
|
||||
# Security: Ensure the path is within the static directory
|
||||
try:
|
||||
widget_path.resolve().relative_to(self.static_dir.resolve())
|
||||
except ValueError:
|
||||
return HTMLResponse(
|
||||
content=f"<h1>403 - Forbidden</h1>",
|
||||
status_code=403
|
||||
)
|
||||
|
||||
logger.info(f"Serving widget: {filename}")
|
||||
return FileResponse(
|
||||
widget_path,
|
||||
media_type="text/html",
|
||||
headers={
|
||||
"Cache-Control": "no-cache, no-store, must-revalidate",
|
||||
"Pragma": "no-cache",
|
||||
"Expires": "0"
|
||||
}
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/widgets",
|
||||
summary="List available widgets"
|
||||
)
|
||||
async def list_widgets():
|
||||
"""
|
||||
List all available widget files
|
||||
|
||||
Returns:
|
||||
List of widget filenames
|
||||
"""
|
||||
widgets_dir = self.static_dir / "widgets"
|
||||
|
||||
if not widgets_dir.exists():
|
||||
return {"widgets": [], "message": "Widgets directory not found"}
|
||||
|
||||
widgets = []
|
||||
for file in widgets_dir.glob("*.html"):
|
||||
widgets.append({
|
||||
"name": file.name,
|
||||
"url": f"/static/widgets/{file.name}",
|
||||
"size": file.stat().st_size
|
||||
})
|
||||
|
||||
return {
|
||||
"widgets": widgets,
|
||||
"count": len(widgets)
|
||||
}
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
static_controller = StaticController()
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
Tools Domain
|
||||
|
||||
Provides utility tool endpoints including DNS lookups and system stats.
|
||||
"""
|
||||
from src.domains.tools.controller import tools_controller
|
||||
from src.domains.tools.dns import DNSService, DNSQueryError
|
||||
from src.domains.tools.system import SystemStatsService, SystemStatsResponse
|
||||
|
||||
__all__ = [
|
||||
"tools_controller",
|
||||
"DNSService",
|
||||
"DNSQueryError",
|
||||
"SystemStatsService",
|
||||
"SystemStatsResponse",
|
||||
]
|
||||
@@ -0,0 +1,153 @@
|
||||
"""
|
||||
Tools Controller
|
||||
|
||||
Provides utility tool endpoints including:
|
||||
- DNS lookups
|
||||
- System stats
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from src.shared.base import BaseController
|
||||
from src.shared.logging import get_logger
|
||||
from src.domains.tools.dns.schemas import DNSLookupRequest, DNSLookupResponse
|
||||
from src.domains.tools.dns.service import DNSService
|
||||
from src.domains.tools.dns.exceptions import DNSQueryError
|
||||
from src.domains.tools.system.schemas import SystemStatsResponse
|
||||
from src.domains.tools.system.service import SystemStatsService
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ToolsController(BaseController):
|
||||
"""
|
||||
Controller for utility tools
|
||||
|
||||
Provides endpoints for:
|
||||
- DNS lookups
|
||||
- System stats
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/tools", tags=["Tools"])
|
||||
self.dns_service = DNSService()
|
||||
self.system_stats_service = SystemStatsService()
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.post(
|
||||
"/dns/lookup",
|
||||
response_model=DNSLookupResponse,
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Perform DNS lookup",
|
||||
description="""
|
||||
Perform DNS lookups for various record types.
|
||||
|
||||
Uses dnspython for reliable DNS queries with support for multiple record types
|
||||
and custom nameservers. Perfect for troubleshooting DNS issues and checking
|
||||
domain configurations.
|
||||
|
||||
**Supported Record Types:**
|
||||
- A: IPv4 address records
|
||||
- AAAA: IPv6 address records
|
||||
- MX: Mail exchange records
|
||||
- TXT: Text records (SPF, DKIM, etc.)
|
||||
- CNAME: Canonical name records
|
||||
- NS: Nameserver records
|
||||
- SOA: Start of authority records
|
||||
- PTR: Pointer records (reverse DNS)
|
||||
- CAA: Certification authority authorization
|
||||
- SRV: Service records
|
||||
|
||||
**Features:**
|
||||
- Custom nameserver support (e.g., 8.8.8.8, 1.1.1.1)
|
||||
- Query time measurement
|
||||
- Detailed error messages
|
||||
|
||||
**Rate Limiting:** None (internal network use only)
|
||||
"""
|
||||
)
|
||||
async def dns_lookup(request: DNSLookupRequest) -> DNSLookupResponse:
|
||||
"""
|
||||
Perform DNS lookup for a domain
|
||||
|
||||
Args:
|
||||
request: DNS lookup request with domain, record type, and optional nameserver
|
||||
|
||||
Returns:
|
||||
DNS lookup results with records and metadata
|
||||
|
||||
Raises:
|
||||
HTTPException: 400 for invalid queries, 500 for processing errors
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Received DNS lookup request for: {request.domain} ({request.record_type})")
|
||||
result = await self.dns_service.lookup(request)
|
||||
return result
|
||||
|
||||
except DNSQueryError as e:
|
||||
logger.warning(f"DNS query error: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"DNS query failed: {str(e)}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error during DNS lookup: {str(e)}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="An unexpected error occurred during DNS lookup"
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/system/stats",
|
||||
response_model=SystemStatsResponse,
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Get host system statistics",
|
||||
description="""
|
||||
Get real-time host system resource statistics.
|
||||
|
||||
Returns CPU, memory, disk, network, and GPU/VRAM usage for the host machine
|
||||
(not Docker container metrics).
|
||||
|
||||
**Metrics Returned:**
|
||||
- **CPU:** Usage percentage, core count, load averages
|
||||
- **Memory:** Usage percentage, total/used/available bytes
|
||||
- **Disk:** Usage percentage, total/used/free bytes (root partition)
|
||||
- **Network:** Total bytes sent/received
|
||||
- **GPU:** VRAM usage (if NVIDIA GPU available via nvidia-smi)
|
||||
|
||||
**Use Cases:**
|
||||
- Dashboard system monitoring widgets
|
||||
- Health checks and alerting
|
||||
- Capacity planning
|
||||
"""
|
||||
)
|
||||
async def get_system_stats() -> SystemStatsResponse:
|
||||
"""
|
||||
Get current host system statistics
|
||||
|
||||
Returns:
|
||||
System statistics including CPU, memory, disk, network, and GPU
|
||||
|
||||
Raises:
|
||||
HTTPException: 500 for processing errors
|
||||
"""
|
||||
try:
|
||||
logger.info("Fetching system stats")
|
||||
result = await self.system_stats_service.get_stats()
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get system stats: {str(e)}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to collect system stats: {str(e)}"
|
||||
)
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
tools_controller = ToolsController()
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
DNS Tools Module
|
||||
|
||||
Provides DNS lookup functionality.
|
||||
"""
|
||||
from src.domains.tools.dns.schemas import DNSLookupRequest, DNSLookupResponse, DNSRecord
|
||||
from src.domains.tools.dns.service import DNSService
|
||||
from src.domains.tools.dns.exceptions import DNSQueryError
|
||||
|
||||
__all__ = [
|
||||
"DNSLookupRequest",
|
||||
"DNSLookupResponse",
|
||||
"DNSRecord",
|
||||
"DNSService",
|
||||
"DNSQueryError",
|
||||
]
|
||||
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
DNS Exceptions
|
||||
"""
|
||||
|
||||
|
||||
class DNSQueryError(Exception):
|
||||
"""Raised when a DNS query fails"""
|
||||
pass
|
||||
@@ -0,0 +1,94 @@
|
||||
"""
|
||||
Pydantic schemas for DNS lookup module
|
||||
"""
|
||||
from pydantic import Field
|
||||
from typing import Optional, List
|
||||
from datetime import datetime
|
||||
from src.shared.base import BaseSchema
|
||||
|
||||
|
||||
class DNSLookupRequest(BaseSchema):
|
||||
"""Request model for DNS lookup"""
|
||||
|
||||
domain: str = Field(
|
||||
...,
|
||||
description="The domain name to lookup",
|
||||
examples=["example.com", "google.com"],
|
||||
min_length=1,
|
||||
max_length=255
|
||||
)
|
||||
|
||||
record_type: str = Field(
|
||||
default="A",
|
||||
description="DNS record type to query (A, AAAA, MX, TXT, CNAME, NS, SOA, PTR, CAA)",
|
||||
examples=["A", "AAAA", "MX", "TXT", "CNAME"]
|
||||
)
|
||||
|
||||
nameserver: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Optional nameserver to use for the query (e.g., 8.8.8.8, 1.1.1.1)",
|
||||
examples=["8.8.8.8", "1.1.1.1", "9.9.9.9"]
|
||||
)
|
||||
|
||||
|
||||
class DNSRecord(BaseSchema):
|
||||
"""Single DNS record result"""
|
||||
|
||||
value: str = Field(
|
||||
...,
|
||||
description="The DNS record value"
|
||||
)
|
||||
|
||||
ttl: Optional[int] = Field(
|
||||
default=None,
|
||||
description="Time to live in seconds"
|
||||
)
|
||||
|
||||
priority: Optional[int] = Field(
|
||||
default=None,
|
||||
description="Priority (for MX records)"
|
||||
)
|
||||
|
||||
|
||||
class DNSLookupResponse(BaseSchema):
|
||||
"""Response model for DNS lookup"""
|
||||
|
||||
domain: str = Field(
|
||||
...,
|
||||
description="The queried domain name"
|
||||
)
|
||||
|
||||
record_type: str = Field(
|
||||
...,
|
||||
description="DNS record type queried"
|
||||
)
|
||||
|
||||
records: List[DNSRecord] = Field(
|
||||
...,
|
||||
description="List of DNS records found"
|
||||
)
|
||||
|
||||
nameserver_used: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Nameserver used for the query"
|
||||
)
|
||||
|
||||
query_time_ms: float = Field(
|
||||
...,
|
||||
description="Query execution time in milliseconds"
|
||||
)
|
||||
|
||||
queried_at: datetime = Field(
|
||||
...,
|
||||
description="UTC timestamp when query was executed"
|
||||
)
|
||||
|
||||
success: bool = Field(
|
||||
...,
|
||||
description="Whether the query was successful"
|
||||
)
|
||||
|
||||
error_message: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Error message if query failed"
|
||||
)
|
||||
@@ -0,0 +1,188 @@
|
||||
"""
|
||||
DNS Lookup Service
|
||||
|
||||
Provides DNS query functionality using dnspython library.
|
||||
"""
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
import dns.resolver
|
||||
import dns.exception
|
||||
|
||||
from src.shared.logging import get_logger
|
||||
from src.domains.tools.dns.schemas import DNSLookupRequest, DNSLookupResponse, DNSRecord
|
||||
from src.domains.tools.dns.exceptions import DNSQueryError
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class DNSService:
|
||||
"""
|
||||
Service for performing DNS lookups
|
||||
|
||||
Uses dnspython for reliable DNS queries with support for
|
||||
various record types and custom nameservers.
|
||||
"""
|
||||
|
||||
SUPPORTED_RECORD_TYPES = [
|
||||
"A", "AAAA", "MX", "TXT", "CNAME", "NS", "SOA", "PTR", "CAA", "SRV"
|
||||
]
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize DNS service"""
|
||||
self.resolver = dns.resolver.Resolver()
|
||||
self.resolver.timeout = 5.0
|
||||
self.resolver.lifetime = 10.0
|
||||
|
||||
async def lookup(self, request: DNSLookupRequest) -> DNSLookupResponse:
|
||||
"""
|
||||
Perform DNS lookup for the specified domain and record type
|
||||
|
||||
Args:
|
||||
request: DNS lookup request with domain, record type, and optional nameserver
|
||||
|
||||
Returns:
|
||||
DNSLookupResponse with query results
|
||||
|
||||
Raises:
|
||||
DNSQueryError: If the DNS query fails
|
||||
"""
|
||||
start_time = time.time()
|
||||
record_type = request.record_type.upper()
|
||||
|
||||
if record_type not in self.SUPPORTED_RECORD_TYPES:
|
||||
raise DNSQueryError(
|
||||
f"Unsupported record type: {record_type}. "
|
||||
f"Supported types: {', '.join(self.SUPPORTED_RECORD_TYPES)}"
|
||||
)
|
||||
|
||||
resolver = dns.resolver.Resolver()
|
||||
resolver.timeout = 5.0
|
||||
resolver.lifetime = 10.0
|
||||
|
||||
nameserver_used = None
|
||||
if request.nameserver:
|
||||
resolver.nameservers = [request.nameserver]
|
||||
nameserver_used = request.nameserver
|
||||
logger.info(f"Using custom nameserver: {request.nameserver}")
|
||||
else:
|
||||
nameserver_used = resolver.nameservers[0] if resolver.nameservers else "system"
|
||||
|
||||
try:
|
||||
logger.info(f"Performing DNS lookup: {request.domain} ({record_type})")
|
||||
|
||||
answers = resolver.resolve(request.domain, record_type)
|
||||
|
||||
records = []
|
||||
for rdata in answers:
|
||||
record = self._parse_record(rdata, record_type)
|
||||
if record:
|
||||
records.append(record)
|
||||
|
||||
query_time_ms = (time.time() - start_time) * 1000
|
||||
|
||||
logger.info(
|
||||
f"DNS lookup successful: {request.domain} ({record_type}) - "
|
||||
f"Found {len(records)} records in {query_time_ms:.2f}ms"
|
||||
)
|
||||
|
||||
return DNSLookupResponse(
|
||||
domain=request.domain,
|
||||
record_type=record_type,
|
||||
records=records,
|
||||
nameserver_used=nameserver_used,
|
||||
query_time_ms=round(query_time_ms, 2),
|
||||
queried_at=datetime.now(timezone.utc),
|
||||
success=True,
|
||||
error_message=None
|
||||
)
|
||||
|
||||
except dns.resolver.NXDOMAIN:
|
||||
error_msg = f"Domain not found: {request.domain}"
|
||||
logger.warning(error_msg)
|
||||
return self._error_response(request, nameserver_used, start_time, error_msg)
|
||||
|
||||
except dns.resolver.NoAnswer:
|
||||
error_msg = f"No {record_type} records found for {request.domain}"
|
||||
logger.warning(error_msg)
|
||||
return self._error_response(request, nameserver_used, start_time, error_msg)
|
||||
|
||||
except dns.resolver.Timeout:
|
||||
error_msg = f"DNS query timeout for {request.domain}"
|
||||
logger.error(error_msg)
|
||||
return self._error_response(request, nameserver_used, start_time, error_msg)
|
||||
|
||||
except dns.exception.DNSException as e:
|
||||
error_msg = f"DNS error: {str(e)}"
|
||||
logger.error(f"DNS query failed for {request.domain}: {e}")
|
||||
return self._error_response(request, nameserver_used, start_time, error_msg)
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"Unexpected error: {str(e)}"
|
||||
logger.error(f"Unexpected error during DNS lookup: {e}", exc_info=True)
|
||||
return self._error_response(request, nameserver_used, start_time, error_msg)
|
||||
|
||||
def _parse_record(self, rdata, record_type: str) -> Optional[DNSRecord]:
|
||||
"""Parse DNS record data into DNSRecord schema"""
|
||||
try:
|
||||
if record_type == "A" or record_type == "AAAA":
|
||||
return DNSRecord(value=str(rdata), ttl=None)
|
||||
|
||||
elif record_type == "MX":
|
||||
return DNSRecord(
|
||||
value=str(rdata.exchange),
|
||||
priority=rdata.preference,
|
||||
ttl=None
|
||||
)
|
||||
|
||||
elif record_type == "TXT":
|
||||
txt_value = " ".join([s.decode() if isinstance(s, bytes) else str(s) for s in rdata.strings])
|
||||
return DNSRecord(value=txt_value, ttl=None)
|
||||
|
||||
elif record_type in ["CNAME", "NS", "PTR"]:
|
||||
return DNSRecord(value=str(rdata.target), ttl=None)
|
||||
|
||||
elif record_type == "SOA":
|
||||
soa_value = f"mname={rdata.mname} rname={rdata.rname} serial={rdata.serial}"
|
||||
return DNSRecord(value=soa_value, ttl=None)
|
||||
|
||||
elif record_type == "CAA":
|
||||
caa_value = f"{rdata.flags} {rdata.tag.decode() if isinstance(rdata.tag, bytes) else rdata.tag} {rdata.value.decode() if isinstance(rdata.value, bytes) else rdata.value}"
|
||||
return DNSRecord(value=caa_value, ttl=None)
|
||||
|
||||
elif record_type == "SRV":
|
||||
srv_value = f"{rdata.target} port={rdata.port} priority={rdata.priority} weight={rdata.weight}"
|
||||
return DNSRecord(
|
||||
value=srv_value,
|
||||
priority=rdata.priority,
|
||||
ttl=None
|
||||
)
|
||||
|
||||
else:
|
||||
return DNSRecord(value=str(rdata), ttl=None)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to parse {record_type} record: {e}")
|
||||
return None
|
||||
|
||||
def _error_response(
|
||||
self,
|
||||
request: DNSLookupRequest,
|
||||
nameserver_used: Optional[str],
|
||||
start_time: float,
|
||||
error_message: str
|
||||
) -> DNSLookupResponse:
|
||||
"""Create an error response for failed DNS queries"""
|
||||
query_time_ms = (time.time() - start_time) * 1000
|
||||
|
||||
return DNSLookupResponse(
|
||||
domain=request.domain,
|
||||
record_type=request.record_type.upper(),
|
||||
records=[],
|
||||
nameserver_used=nameserver_used,
|
||||
query_time_ms=round(query_time_ms, 2),
|
||||
queried_at=datetime.now(timezone.utc),
|
||||
success=False,
|
||||
error_message=error_message
|
||||
)
|
||||
@@ -0,0 +1,6 @@
|
||||
"""System stats module for host system resource monitoring."""
|
||||
|
||||
from src.domains.tools.system.service import SystemStatsService
|
||||
from src.domains.tools.system.schemas import SystemStatsResponse
|
||||
|
||||
__all__ = ["SystemStatsService", "SystemStatsResponse"]
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
Pydantic schemas for system stats module
|
||||
"""
|
||||
from pydantic import Field
|
||||
from typing import Optional, List
|
||||
from datetime import datetime
|
||||
from src.shared.base import BaseSchema
|
||||
|
||||
|
||||
class CpuStats(BaseSchema):
|
||||
"""CPU usage statistics"""
|
||||
|
||||
usage_percent: float = Field(
|
||||
...,
|
||||
description="CPU usage percentage (0-100)",
|
||||
ge=0,
|
||||
le=100
|
||||
)
|
||||
|
||||
cores: int = Field(
|
||||
...,
|
||||
description="Number of CPU cores"
|
||||
)
|
||||
|
||||
load_1m: Optional[float] = Field(
|
||||
default=None,
|
||||
description="1-minute load average"
|
||||
)
|
||||
|
||||
load_5m: Optional[float] = Field(
|
||||
default=None,
|
||||
description="5-minute load average"
|
||||
)
|
||||
|
||||
load_15m: Optional[float] = Field(
|
||||
default=None,
|
||||
description="15-minute load average"
|
||||
)
|
||||
|
||||
|
||||
class MemoryStats(BaseSchema):
|
||||
"""Memory usage statistics"""
|
||||
|
||||
usage_percent: float = Field(
|
||||
...,
|
||||
description="Memory usage percentage (0-100)",
|
||||
ge=0,
|
||||
le=100
|
||||
)
|
||||
|
||||
total_bytes: int = Field(
|
||||
...,
|
||||
description="Total memory in bytes"
|
||||
)
|
||||
|
||||
used_bytes: int = Field(
|
||||
...,
|
||||
description="Used memory in bytes"
|
||||
)
|
||||
|
||||
available_bytes: int = Field(
|
||||
...,
|
||||
description="Available memory in bytes"
|
||||
)
|
||||
|
||||
|
||||
class DiskStats(BaseSchema):
|
||||
"""Disk usage statistics for a single mount point"""
|
||||
|
||||
mount_point: str = Field(
|
||||
...,
|
||||
description="Mount point path"
|
||||
)
|
||||
|
||||
device: str = Field(
|
||||
...,
|
||||
description="Device name (e.g., /dev/sda1)"
|
||||
)
|
||||
|
||||
fstype: str = Field(
|
||||
...,
|
||||
description="Filesystem type (e.g., ext4, xfs)"
|
||||
)
|
||||
|
||||
usage_percent: float = Field(
|
||||
...,
|
||||
description="Disk usage percentage (0-100)",
|
||||
ge=0,
|
||||
le=100
|
||||
)
|
||||
|
||||
total_bytes: int = Field(
|
||||
...,
|
||||
description="Total disk space in bytes"
|
||||
)
|
||||
|
||||
used_bytes: int = Field(
|
||||
...,
|
||||
description="Used disk space in bytes"
|
||||
)
|
||||
|
||||
free_bytes: int = Field(
|
||||
...,
|
||||
description="Free disk space in bytes"
|
||||
)
|
||||
|
||||
|
||||
class NetworkStats(BaseSchema):
|
||||
"""Network I/O statistics"""
|
||||
|
||||
bytes_sent: int = Field(
|
||||
...,
|
||||
description="Total bytes sent"
|
||||
)
|
||||
|
||||
bytes_recv: int = Field(
|
||||
...,
|
||||
description="Total bytes received"
|
||||
)
|
||||
|
||||
bytes_total: int = Field(
|
||||
...,
|
||||
description="Total bytes (sent + received)"
|
||||
)
|
||||
|
||||
|
||||
class GpuStats(BaseSchema):
|
||||
"""GPU/VRAM statistics (if available)"""
|
||||
|
||||
available: bool = Field(
|
||||
...,
|
||||
description="Whether GPU stats are available"
|
||||
)
|
||||
|
||||
name: Optional[str] = Field(
|
||||
default=None,
|
||||
description="GPU name"
|
||||
)
|
||||
|
||||
usage_percent: Optional[float] = Field(
|
||||
default=None,
|
||||
description="VRAM usage percentage (0-100)"
|
||||
)
|
||||
|
||||
total_bytes: Optional[int] = Field(
|
||||
default=None,
|
||||
description="Total VRAM in bytes"
|
||||
)
|
||||
|
||||
used_bytes: Optional[int] = Field(
|
||||
default=None,
|
||||
description="Used VRAM in bytes"
|
||||
)
|
||||
|
||||
free_bytes: Optional[int] = Field(
|
||||
default=None,
|
||||
description="Free VRAM in bytes"
|
||||
)
|
||||
|
||||
|
||||
class SystemStatsResponse(BaseSchema):
|
||||
"""Response model for system stats"""
|
||||
|
||||
cpu: CpuStats = Field(
|
||||
...,
|
||||
description="CPU statistics"
|
||||
)
|
||||
|
||||
memory: MemoryStats = Field(
|
||||
...,
|
||||
description="Memory statistics"
|
||||
)
|
||||
|
||||
disks: List[DiskStats] = Field(
|
||||
...,
|
||||
description="Disk statistics for all mounted filesystems"
|
||||
)
|
||||
|
||||
network: NetworkStats = Field(
|
||||
...,
|
||||
description="Network I/O statistics"
|
||||
)
|
||||
|
||||
gpu: GpuStats = Field(
|
||||
...,
|
||||
description="GPU/VRAM statistics"
|
||||
)
|
||||
|
||||
hostname: str = Field(
|
||||
...,
|
||||
description="System hostname"
|
||||
)
|
||||
|
||||
queried_at: datetime = Field(
|
||||
...,
|
||||
description="UTC timestamp when stats were collected"
|
||||
)
|
||||
@@ -0,0 +1,190 @@
|
||||
"""
|
||||
System stats service for collecting host system metrics
|
||||
"""
|
||||
import subprocess
|
||||
import socket
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import psutil
|
||||
|
||||
from src.shared.logging import get_logger
|
||||
from typing import List
|
||||
|
||||
from src.domains.tools.system.schemas import (
|
||||
SystemStatsResponse,
|
||||
CpuStats,
|
||||
MemoryStats,
|
||||
DiskStats,
|
||||
NetworkStats,
|
||||
GpuStats,
|
||||
)
|
||||
|
||||
# Filesystem types to exclude (virtual/system filesystems)
|
||||
EXCLUDED_FSTYPES = {
|
||||
"tmpfs", "devtmpfs", "devfs", "squashfs", "overlay",
|
||||
"aufs", "proc", "sysfs", "cgroup", "cgroup2",
|
||||
"debugfs", "tracefs", "securityfs", "pstore",
|
||||
"hugetlbfs", "mqueue", "binfmt_misc", "autofs",
|
||||
"fuse.lxcfs", "nsfs", "efivarfs",
|
||||
}
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class SystemStatsService:
|
||||
"""Service for collecting host system statistics"""
|
||||
|
||||
async def get_stats(self) -> SystemStatsResponse:
|
||||
"""
|
||||
Collect current system statistics.
|
||||
|
||||
Returns:
|
||||
SystemStatsResponse with CPU, memory, disks, network, and GPU stats
|
||||
"""
|
||||
cpu = self._get_cpu_stats()
|
||||
memory = self._get_memory_stats()
|
||||
disks = self._get_all_disk_stats()
|
||||
network = self._get_network_stats()
|
||||
gpu = self._get_gpu_stats()
|
||||
|
||||
return SystemStatsResponse(
|
||||
cpu=cpu,
|
||||
memory=memory,
|
||||
disks=disks,
|
||||
network=network,
|
||||
gpu=gpu,
|
||||
hostname=socket.gethostname(),
|
||||
queried_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
def _get_cpu_stats(self) -> CpuStats:
|
||||
"""Get CPU usage statistics"""
|
||||
# Get CPU percentage (blocking call with interval for accuracy)
|
||||
cpu_percent = psutil.cpu_percent(interval=0.1)
|
||||
cpu_count = psutil.cpu_count()
|
||||
|
||||
# Get load averages (Unix only)
|
||||
try:
|
||||
load_1, load_5, load_15 = psutil.getloadavg()
|
||||
except (AttributeError, OSError):
|
||||
load_1 = load_5 = load_15 = None
|
||||
|
||||
return CpuStats(
|
||||
usage_percent=cpu_percent,
|
||||
cores=cpu_count or 1,
|
||||
load_1m=load_1,
|
||||
load_5m=load_5,
|
||||
load_15m=load_15,
|
||||
)
|
||||
|
||||
def _get_memory_stats(self) -> MemoryStats:
|
||||
"""Get memory usage statistics"""
|
||||
mem = psutil.virtual_memory()
|
||||
|
||||
return MemoryStats(
|
||||
usage_percent=mem.percent,
|
||||
total_bytes=mem.total,
|
||||
used_bytes=mem.used,
|
||||
available_bytes=mem.available,
|
||||
)
|
||||
|
||||
def _get_all_disk_stats(self) -> List[DiskStats]:
|
||||
"""Get disk usage statistics for all mounted real filesystems"""
|
||||
disks = []
|
||||
seen_devices = set()
|
||||
|
||||
for partition in psutil.disk_partitions(all=False):
|
||||
# Skip excluded filesystem types
|
||||
if partition.fstype.lower() in EXCLUDED_FSTYPES:
|
||||
continue
|
||||
|
||||
# Skip duplicate devices (same device mounted multiple times)
|
||||
if partition.device in seen_devices:
|
||||
continue
|
||||
seen_devices.add(partition.device)
|
||||
|
||||
# Skip Docker/container overlays
|
||||
if partition.mountpoint.startswith("/var/lib/docker"):
|
||||
continue
|
||||
|
||||
try:
|
||||
usage = psutil.disk_usage(partition.mountpoint)
|
||||
disks.append(DiskStats(
|
||||
mount_point=partition.mountpoint,
|
||||
device=partition.device,
|
||||
fstype=partition.fstype,
|
||||
usage_percent=usage.percent,
|
||||
total_bytes=usage.total,
|
||||
used_bytes=usage.used,
|
||||
free_bytes=usage.free,
|
||||
))
|
||||
except (PermissionError, OSError) as e:
|
||||
logger.debug(f"Skipping {partition.mountpoint}: {e}")
|
||||
continue
|
||||
|
||||
# Sort by mount point for consistent ordering
|
||||
disks.sort(key=lambda d: d.mount_point)
|
||||
|
||||
return disks
|
||||
|
||||
def _get_network_stats(self) -> NetworkStats:
|
||||
"""Get network I/O statistics"""
|
||||
net_io = psutil.net_io_counters()
|
||||
|
||||
return NetworkStats(
|
||||
bytes_sent=net_io.bytes_sent,
|
||||
bytes_recv=net_io.bytes_recv,
|
||||
bytes_total=net_io.bytes_sent + net_io.bytes_recv,
|
||||
)
|
||||
|
||||
def _get_gpu_stats(self) -> GpuStats:
|
||||
"""Get GPU/VRAM statistics using nvidia-smi"""
|
||||
try:
|
||||
# Query nvidia-smi for GPU memory info
|
||||
result = subprocess.run(
|
||||
[
|
||||
"nvidia-smi",
|
||||
"--query-gpu=name,memory.total,memory.used,memory.free",
|
||||
"--format=csv,noheader,nounits",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.debug("nvidia-smi not available or failed")
|
||||
return GpuStats(available=False)
|
||||
|
||||
# Parse output: "NVIDIA GeForce RTX 3080, 10240, 2048, 8192"
|
||||
line = result.stdout.strip().split("\n")[0] # First GPU
|
||||
parts = [p.strip() for p in line.split(",")]
|
||||
|
||||
if len(parts) >= 4:
|
||||
name = parts[0]
|
||||
total_mb = int(parts[1])
|
||||
used_mb = int(parts[2])
|
||||
free_mb = int(parts[3])
|
||||
|
||||
total_bytes = total_mb * 1024 * 1024
|
||||
used_bytes = used_mb * 1024 * 1024
|
||||
free_bytes = free_mb * 1024 * 1024
|
||||
usage_percent = (used_mb / total_mb * 100) if total_mb > 0 else 0
|
||||
|
||||
return GpuStats(
|
||||
available=True,
|
||||
name=name,
|
||||
usage_percent=round(usage_percent, 1),
|
||||
total_bytes=total_bytes,
|
||||
used_bytes=used_bytes,
|
||||
free_bytes=free_bytes,
|
||||
)
|
||||
|
||||
except FileNotFoundError:
|
||||
logger.debug("nvidia-smi not found - no NVIDIA GPU available")
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.warning("nvidia-smi timed out")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get GPU stats: {e}")
|
||||
|
||||
return GpuStats(available=False)
|
||||
+37
-72
@@ -6,15 +6,20 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from src.config import get_settings
|
||||
from src.logging_config import setup_logging, get_logger
|
||||
from src.shared.config import get_settings
|
||||
from src.shared.logging import setup_logging, get_logger
|
||||
from src.shared.database import get_database
|
||||
from src.shared.security import initialize_oidc
|
||||
from src.models.ollama_client import get_ollama_client, close_ollama_client
|
||||
from src.controllers.infrastructure_controller import infrastructure_controller
|
||||
from src.controllers.tools_controller import tools_controller
|
||||
from src.controllers.health_controller import health_controller
|
||||
from src.controllers.static_controller import static_controller
|
||||
from src.controllers.ai_controller import router as ai_router
|
||||
from src.security import initialize_oidc
|
||||
|
||||
# Import domain controllers
|
||||
from src.domains.health import health_controller
|
||||
from src.domains.auth import auth_controller
|
||||
from src.domains.tools import tools_controller
|
||||
from src.domains.infrastructure import infrastructure_controller
|
||||
from src.domains.housekeeping import housekeeping_controller
|
||||
from src.domains.static import static_controller
|
||||
from src.domains.dashboard import dashboard_controller
|
||||
|
||||
# Initialize settings
|
||||
settings = get_settings()
|
||||
@@ -44,9 +49,17 @@ async def lifespan(app: FastAPI):
|
||||
ollama_client = get_ollama_client()
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
if ollama_healthy:
|
||||
logger.info("✓ Ollama connection successful")
|
||||
logger.info("Ollama connection successful")
|
||||
else:
|
||||
logger.warning("✗ Ollama connection failed - AI features may not work")
|
||||
logger.warning("Ollama connection failed - AI features may not work")
|
||||
|
||||
# Check database connectivity
|
||||
database = get_database()
|
||||
db_healthy = await database.health_check()
|
||||
if db_healthy:
|
||||
logger.info("Database connection successful")
|
||||
else:
|
||||
logger.warning("Database connection failed - auth features may not work")
|
||||
|
||||
# Initialize security (OIDC authentication)
|
||||
initialize_oidc(settings)
|
||||
@@ -56,6 +69,7 @@ async def lifespan(app: FastAPI):
|
||||
# Shutdown
|
||||
logger.info("Shutting down application")
|
||||
await close_ollama_client()
|
||||
await database.close()
|
||||
|
||||
|
||||
# Create FastAPI application
|
||||
@@ -63,70 +77,19 @@ app = FastAPI(
|
||||
title=settings.app_name,
|
||||
version=settings.app_version,
|
||||
description="""
|
||||
Core Code API provides OpenAPI-compatible functions and AI orchestration for Open WebUI.
|
||||
Core Code API - Infrastructure management and home automation API.
|
||||
|
||||
## Features
|
||||
## Features
|
||||
|
||||
### OpenAI-Compatible API (v1)
|
||||
- `/v1/chat/completions` - Chat completions with streaming support
|
||||
- `/v1/models` - List available models
|
||||
Compatible with OpenAI client libraries and Open WebUI.
|
||||
- **Infrastructure Management** - Container and stack management via Portainer
|
||||
- **Home Automation** - Device control via Home Assistant
|
||||
- **Dashboard** - Quick links and widget management
|
||||
- **Tools** - DNS lookup and utilities
|
||||
|
||||
### Conversation Memory (Phase 2)
|
||||
- `/v1/conversations/{id}` - Get conversation history
|
||||
- `/v1/conversations/{id}/search` - Semantic search within conversation
|
||||
- `/v1/conversations/search` - Search across all conversations
|
||||
- `/v1/conversations/{id}/stats` - Get conversation statistics
|
||||
- `/v1/conversations/{id}/consolidate` - Manual consolidation
|
||||
- `DELETE /v1/conversations/{id}` - Delete conversation
|
||||
|
||||
Multi-tier memory system:
|
||||
- **Tier 1**: Fast in-memory buffer (last 10 turns)
|
||||
- **Tier 2/3**: Unified Qdrant storage (persistent + semantic search)
|
||||
|
||||
### Infrastructure Management
|
||||
**Read Endpoints:**
|
||||
- `GET /infrastructure/health` - Check Portainer & NPM connectivity
|
||||
- `GET /infrastructure/services` - List all deployed services
|
||||
- `GET /infrastructure/services/{name}` - Get service details
|
||||
- `GET /infrastructure/ports` - List allocated ports
|
||||
- `GET /infrastructure/domains` - List configured domains
|
||||
|
||||
**Write Endpoints (Admin Only):**
|
||||
- `POST /infrastructure/services` - Deploy new service from compose YAML
|
||||
- `PUT /infrastructure/services/{name}` - Update existing service
|
||||
- `DELETE /infrastructure/services/{name}` - Remove service and stack
|
||||
- `POST /infrastructure/proxy` - Create proxy host with optional SSL
|
||||
|
||||
Automates infrastructure operations via Portainer and Nginx Proxy Manager APIs.
|
||||
|
||||
### Web Scraper
|
||||
Intelligent web scraping with main content extraction.
|
||||
Perfect for extracting articles, documentation, and blog posts for LLM consumption.
|
||||
|
||||
## Authentication
|
||||
|
||||
When OIDC authentication is enabled (oidc_enabled=true in config):
|
||||
- Infrastructure write endpoints require authentication
|
||||
- Use OAuth2/OIDC bearer token from Authentik
|
||||
- Admin group membership required for infrastructure operations
|
||||
|
||||
## Integration
|
||||
|
||||
This API is designed to integrate with:
|
||||
- **Open WebUI**: Direct OpenAI API compatibility
|
||||
- **Open WebUI Functions**: Import via OpenAPI spec
|
||||
- **Open WebUI Pipelines**: Use as data source
|
||||
- **LangChain**: Compatible with standard HTTP tools
|
||||
|
||||
## Documentation
|
||||
|
||||
- **OpenAPI Spec**: `/openapi.json`
|
||||
- **Swagger UI**: `/docs`
|
||||
- **ReDoc**: `/redoc`
|
||||
See `/docs` for the full API reference.
|
||||
""",
|
||||
docs_url="/docs",
|
||||
redoc_url="/redoc",
|
||||
redoc_url=None,
|
||||
openapi_url="/openapi.json",
|
||||
lifespan=lifespan,
|
||||
debug=settings.debug,
|
||||
@@ -146,12 +109,14 @@ app.add_middleware(
|
||||
)
|
||||
|
||||
|
||||
# Include controller routers
|
||||
# Include domain routers
|
||||
app.include_router(health_controller.router) # / and /health
|
||||
app.include_router(tools_controller.router) # /web-scraper/scrape
|
||||
app.include_router(auth_controller.router) # /auth/*
|
||||
app.include_router(tools_controller.router) # /tools/*
|
||||
app.include_router(infrastructure_controller.router) # /infrastructure/*
|
||||
app.include_router(housekeeping_controller.router) # /housekeeping/*
|
||||
app.include_router(static_controller.router) # /static/*
|
||||
app.include_router(ai_router) # /ai/*
|
||||
app.include_router(dashboard_controller.router) # /dashboard/*
|
||||
|
||||
|
||||
# Global exception handler
|
||||
|
||||
@@ -10,7 +10,6 @@ ALWAYS_ON_SERVICES: Set[str] = {
|
||||
"portainer",
|
||||
"nginx-proxy-manager",
|
||||
"core-api",
|
||||
"uptime-kuma",
|
||||
"organizr",
|
||||
"headscale",
|
||||
"watchtower",
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
"""
|
||||
Shared utilities for Core-API
|
||||
|
||||
Contains common base classes, configuration, database, logging utilities,
|
||||
and API clients used across all domains.
|
||||
"""
|
||||
from src.shared.config import get_settings, Settings
|
||||
from src.shared.database import Base, get_database, get_async_session
|
||||
from src.shared.logging import get_logger, setup_logging
|
||||
from src.shared.base import BaseController, BaseSchema
|
||||
from src.shared.security import initialize_oidc
|
||||
|
||||
# Re-export clients for convenience
|
||||
from src.shared.clients import (
|
||||
PortainerClient,
|
||||
get_portainer_client,
|
||||
NPMClient,
|
||||
get_npm_client,
|
||||
HomeAssistantClient,
|
||||
get_homeassistant_client,
|
||||
AuthentikClient,
|
||||
get_authentik_client,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Config
|
||||
"get_settings",
|
||||
"Settings",
|
||||
# Database
|
||||
"Base",
|
||||
"get_database",
|
||||
"get_async_session",
|
||||
# Logging
|
||||
"get_logger",
|
||||
"setup_logging",
|
||||
# Base classes
|
||||
"BaseController",
|
||||
"BaseSchema",
|
||||
# Security
|
||||
"initialize_oidc",
|
||||
# Clients
|
||||
"PortainerClient",
|
||||
"get_portainer_client",
|
||||
"NPMClient",
|
||||
"get_npm_client",
|
||||
"HomeAssistantClient",
|
||||
"get_homeassistant_client",
|
||||
"AuthentikClient",
|
||||
"get_authentik_client",
|
||||
]
|
||||
@@ -0,0 +1,65 @@
|
||||
"""
|
||||
Base classes for Core-API
|
||||
|
||||
Provides common base classes for controllers and schemas.
|
||||
"""
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from abc import ABC, abstractmethod
|
||||
from fastapi import APIRouter
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
|
||||
class BaseController(ABC):
|
||||
"""
|
||||
Base controller class with common functionality
|
||||
|
||||
All controllers should inherit from this class and implement
|
||||
the create_router() method to define their endpoints.
|
||||
"""
|
||||
|
||||
def __init__(self, prefix: str, tags: list[str]):
|
||||
"""
|
||||
Initialize base controller
|
||||
|
||||
Args:
|
||||
prefix: URL prefix for this controller's routes
|
||||
tags: OpenAPI tags for documentation grouping
|
||||
"""
|
||||
self.prefix = prefix
|
||||
self.tags = tags
|
||||
self._router = None
|
||||
|
||||
@abstractmethod
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the FastAPI router for this controller"""
|
||||
pass
|
||||
|
||||
@property
|
||||
def router(self) -> APIRouter:
|
||||
"""Get the router instance, creating it if needed"""
|
||||
if self._router is None:
|
||||
self._router = self.create_router()
|
||||
return self._router
|
||||
|
||||
|
||||
class BaseSchema(BaseModel):
|
||||
"""
|
||||
Base Pydantic model with standardized configuration
|
||||
|
||||
All schemas should inherit from this to ensure consistent behavior.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
strict=False,
|
||||
populate_by_name=True,
|
||||
use_enum_values=True,
|
||||
validate_assignment=True,
|
||||
json_encoders={
|
||||
datetime: lambda v: v.isoformat() if v else None
|
||||
}
|
||||
)
|
||||
|
||||
def dict_without_none(self) -> dict[str, Any]:
|
||||
"""Return model as dict, excluding None values"""
|
||||
return {k: v for k, v in self.model_dump().items() if v is not None}
|
||||
@@ -0,0 +1,20 @@
|
||||
"""
|
||||
API Clients for Core-API
|
||||
|
||||
Provides HTTP/WebSocket clients for external infrastructure services.
|
||||
"""
|
||||
from src.shared.clients.portainer_client import PortainerClient, get_portainer_client
|
||||
from src.shared.clients.npm_client import NPMClient, get_npm_client
|
||||
from src.shared.clients.homeassistant_client import HomeAssistantClient, get_homeassistant_client
|
||||
from src.shared.clients.authentik_client import AuthentikClient, get_authentik_client
|
||||
|
||||
__all__ = [
|
||||
"PortainerClient",
|
||||
"get_portainer_client",
|
||||
"NPMClient",
|
||||
"get_npm_client",
|
||||
"HomeAssistantClient",
|
||||
"get_homeassistant_client",
|
||||
"AuthentikClient",
|
||||
"get_authentik_client",
|
||||
]
|
||||
@@ -0,0 +1,302 @@
|
||||
"""
|
||||
Authentik API Client
|
||||
|
||||
Provides methods for interacting with Authentik Identity Provider API.
|
||||
Used for managing applications, providers, and authentication flows.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Dict, List, Any, Optional
|
||||
from functools import lru_cache
|
||||
from src.shared.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AuthentikClient:
|
||||
"""Client for Authentik API operations"""
|
||||
|
||||
def __init__(self, base_url: str, api_token: str):
|
||||
"""
|
||||
Initialize Authentik client
|
||||
|
||||
Args:
|
||||
base_url: Authentik base URL (e.g., http://authentik-server:9000)
|
||||
api_token: API token for authentication
|
||||
"""
|
||||
self.base_url = base_url.rstrip('/')
|
||||
self.api_token = api_token
|
||||
self.client = httpx.AsyncClient(timeout=30.0)
|
||||
|
||||
async def _request(self, method: str, endpoint: str, **kwargs) -> Dict:
|
||||
"""Make authenticated API request using token auth"""
|
||||
headers = kwargs.pop("headers", {})
|
||||
headers["Authorization"] = f"Bearer {self.api_token}"
|
||||
|
||||
response = await self.client.request(
|
||||
method,
|
||||
f"{self.base_url}/api/v3/{endpoint.lstrip('/')}",
|
||||
headers=headers,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if not response.is_success:
|
||||
logger.error(f"API request failed: {response.status_code}")
|
||||
logger.error(f"Response body: {response.text}")
|
||||
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""Check if Authentik is accessible"""
|
||||
try:
|
||||
response = await self.client.get(f"{self.base_url}/-/health/live/")
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"Authentik health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def create_oauth2_provider(
|
||||
self,
|
||||
name: str,
|
||||
client_id: str,
|
||||
redirect_uris: List[str],
|
||||
authorization_flow_slug: str = "default-provider-authorization-implicit-consent",
|
||||
signing_key: Optional[str] = None
|
||||
) -> Dict:
|
||||
"""
|
||||
Create an OAuth2/OIDC provider
|
||||
|
||||
Args:
|
||||
name: Provider name
|
||||
client_id: OAuth2 client ID
|
||||
redirect_uris: List of allowed redirect URIs
|
||||
authorization_flow_slug: Authorization flow slug (will be resolved to UUID)
|
||||
signing_key: Signing key UUID (defaults to auto-selected)
|
||||
|
||||
Returns:
|
||||
Created provider data including client_secret
|
||||
"""
|
||||
# Get authorization flow UUID from slug
|
||||
flows = await self.list_flows()
|
||||
auth_flow_uuid = None
|
||||
invalidation_flow_uuid = None
|
||||
|
||||
for flow in flows:
|
||||
if flow.get("slug") == authorization_flow_slug:
|
||||
auth_flow_uuid = flow.get("pk")
|
||||
if flow.get("slug") == "default-provider-invalidation-flow":
|
||||
invalidation_flow_uuid = flow.get("pk")
|
||||
|
||||
if not auth_flow_uuid:
|
||||
raise ValueError(f"Authorization flow '{authorization_flow_slug}' not found")
|
||||
if not invalidation_flow_uuid:
|
||||
raise ValueError("Invalidation flow not found")
|
||||
|
||||
# Get signing key if not provided
|
||||
if not signing_key:
|
||||
keys = await self._request("GET", "crypto/certificatekeypairs/")
|
||||
# Find the self-signed cert
|
||||
for key in keys.get("results", []):
|
||||
if "authentik" in key.get("name", "").lower():
|
||||
signing_key = key.get("pk")
|
||||
break
|
||||
|
||||
if not signing_key and keys.get("results"):
|
||||
signing_key = keys["results"][0]["pk"]
|
||||
|
||||
# Format redirect URIs as objects with matching_mode
|
||||
formatted_redirect_uris = [
|
||||
{"url": uri, "matching_mode": "strict"}
|
||||
for uri in redirect_uris
|
||||
]
|
||||
|
||||
provider_data = {
|
||||
"name": name,
|
||||
"authorization_flow": auth_flow_uuid,
|
||||
"invalidation_flow": invalidation_flow_uuid,
|
||||
"client_type": "confidential",
|
||||
"client_id": client_id,
|
||||
"redirect_uris": formatted_redirect_uris,
|
||||
"signing_key": signing_key,
|
||||
"sub_mode": "hashed_user_id",
|
||||
"include_claims_in_id_token": True,
|
||||
"issuer_mode": "per_provider",
|
||||
"access_token_validity": "minutes=60",
|
||||
"refresh_token_validity": "days=30",
|
||||
"property_mappings": [] # Will use default mappings
|
||||
}
|
||||
|
||||
result = await self._request("POST", "providers/oauth2/", json=provider_data)
|
||||
logger.info(f"Created OAuth2 provider: {name} (ID: {result.get('pk')})")
|
||||
return result
|
||||
|
||||
async def create_application(
|
||||
self,
|
||||
name: str,
|
||||
slug: str,
|
||||
provider_pk: int,
|
||||
launch_url: Optional[str] = None,
|
||||
icon_url: Optional[str] = None
|
||||
) -> Dict:
|
||||
"""
|
||||
Create an application
|
||||
|
||||
Args:
|
||||
name: Application display name
|
||||
slug: Application slug (URL-safe identifier)
|
||||
provider_pk: Primary key of the provider to use
|
||||
launch_url: Optional launch URL
|
||||
icon_url: Optional icon URL
|
||||
|
||||
Returns:
|
||||
Created application data
|
||||
"""
|
||||
app_data = {
|
||||
"name": name,
|
||||
"slug": slug,
|
||||
"provider": provider_pk,
|
||||
"meta_launch_url": launch_url or "",
|
||||
"meta_icon": icon_url or "",
|
||||
"policy_engine_mode": "any",
|
||||
"open_in_new_tab": False
|
||||
}
|
||||
|
||||
result = await self._request("POST", "core/applications/", json=app_data)
|
||||
logger.info(f"Created application: {name} (slug: {slug})")
|
||||
return result
|
||||
|
||||
async def get_provider_by_name(self, name: str) -> Optional[Dict]:
|
||||
"""Get OAuth2 provider by name"""
|
||||
providers = await self._request("GET", "providers/oauth2/", params={"name": name})
|
||||
results = providers.get("results", [])
|
||||
return results[0] if results else None
|
||||
|
||||
async def get_application_by_slug(self, slug: str) -> Optional[Dict]:
|
||||
"""Get application by slug"""
|
||||
apps = await self._request("GET", "core/applications/", params={"slug": slug})
|
||||
results = apps.get("results", [])
|
||||
return results[0] if results else None
|
||||
|
||||
async def list_flows(self) -> List[Dict]:
|
||||
"""List all authentication flows"""
|
||||
result = await self._request("GET", "flows/instances/")
|
||||
return result.get("results", [])
|
||||
|
||||
async def create_proxy_provider(
|
||||
self,
|
||||
name: str,
|
||||
external_host: str,
|
||||
authorization_flow_slug: str = "default-provider-authorization-implicit-consent",
|
||||
mode: str = "forward_single",
|
||||
token_validity: int = 480 # 8 hours in minutes
|
||||
) -> Dict:
|
||||
"""
|
||||
Create a Proxy Provider for forward authentication
|
||||
|
||||
Args:
|
||||
name: Provider name
|
||||
external_host: External URL (e.g., https://auth.schweitz.net)
|
||||
authorization_flow_slug: Authorization flow slug
|
||||
mode: Proxy mode (forward_single for forward auth)
|
||||
token_validity: Token validity in minutes (default: 480 = 8 hours)
|
||||
|
||||
Returns:
|
||||
Created provider data
|
||||
"""
|
||||
# Get authorization flow UUID from slug
|
||||
flows = await self.list_flows()
|
||||
auth_flow_uuid = None
|
||||
invalidation_flow_uuid = None
|
||||
|
||||
for flow in flows:
|
||||
if flow.get("slug") == authorization_flow_slug:
|
||||
auth_flow_uuid = flow.get("pk")
|
||||
if flow.get("slug") == "default-provider-invalidation-flow":
|
||||
invalidation_flow_uuid = flow.get("pk")
|
||||
|
||||
if not auth_flow_uuid:
|
||||
raise ValueError(f"Authorization flow '{authorization_flow_slug}' not found")
|
||||
if not invalidation_flow_uuid:
|
||||
raise ValueError("Invalidation flow not found")
|
||||
|
||||
provider_data = {
|
||||
"name": name,
|
||||
"authorization_flow": auth_flow_uuid,
|
||||
"invalidation_flow": invalidation_flow_uuid,
|
||||
"mode": mode,
|
||||
"external_host": external_host,
|
||||
"access_token_validity": f"minutes={token_validity}",
|
||||
"refresh_token_validity": f"minutes={token_validity}",
|
||||
"session_duration": f"seconds={token_validity * 60}",
|
||||
"cookie_domain": "", # Will use the domain of each proxied site
|
||||
"property_mappings": []
|
||||
}
|
||||
|
||||
result = await self._request("POST", "providers/proxy/", json=provider_data)
|
||||
logger.info(f"Created Proxy provider: {name} (ID: {result.get('pk')})")
|
||||
return result
|
||||
|
||||
async def get_provider_by_name_proxy(self, name: str) -> Optional[Dict]:
|
||||
"""Get Proxy provider by name"""
|
||||
providers = await self._request("GET", "providers/proxy/", params={"name": name})
|
||||
results = providers.get("results", [])
|
||||
return results[0] if results else None
|
||||
|
||||
async def create_outpost(
|
||||
self,
|
||||
name: str,
|
||||
type: str,
|
||||
providers: List[int],
|
||||
config: Optional[Dict] = None
|
||||
) -> Dict:
|
||||
"""
|
||||
Create an Authentik Outpost
|
||||
|
||||
Args:
|
||||
name: Outpost name
|
||||
type: Outpost type (e.g., "proxy")
|
||||
providers: List of provider PKs
|
||||
config: Optional configuration overrides
|
||||
|
||||
Returns:
|
||||
Created outpost data
|
||||
"""
|
||||
outpost_data = {
|
||||
"name": name,
|
||||
"type": type,
|
||||
"providers": providers,
|
||||
"config": config or {},
|
||||
"service_connection": None # Will use local Docker
|
||||
}
|
||||
|
||||
result = await self._request("POST", "outposts/instances/", json=outpost_data)
|
||||
logger.info(f"Created outpost: {name} (ID: {result.get('pk')})")
|
||||
return result
|
||||
|
||||
async def get_outpost_by_name(self, name: str) -> Optional[Dict]:
|
||||
"""Get outpost by name"""
|
||||
outposts = await self._request("GET", "outposts/instances/", params={"name": name})
|
||||
results = outposts.get("results", [])
|
||||
return results[0] if results else None
|
||||
|
||||
async def close(self):
|
||||
"""Close HTTP client"""
|
||||
await self.client.aclose()
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_authentik_client() -> AuthentikClient:
|
||||
"""Get cached Authentik client instance"""
|
||||
# Import credentials from gitignored module
|
||||
try:
|
||||
from src.credentials import AUTHENTIK_URL, AUTHENTIK_CORE_API_TOKEN
|
||||
except ImportError:
|
||||
# Fallback to environment variables if credentials.py doesn't exist
|
||||
import os
|
||||
AUTHENTIK_URL = os.getenv("AUTHENTIK_URL", "http://authentik-server:9000")
|
||||
AUTHENTIK_CORE_API_TOKEN = os.getenv("AUTHENTIK_API_TOKEN", "")
|
||||
|
||||
return AuthentikClient(
|
||||
base_url=AUTHENTIK_URL,
|
||||
api_token=AUTHENTIK_CORE_API_TOKEN
|
||||
)
|
||||
@@ -0,0 +1,409 @@
|
||||
"""
|
||||
Home Assistant REST API Client
|
||||
|
||||
Provides interface to Home Assistant REST API for home automation control.
|
||||
Uses long-lived access token authentication.
|
||||
API Reference: https://developers.home-assistant.io/docs/api/rest/
|
||||
"""
|
||||
import httpx
|
||||
import json
|
||||
from typing import Optional, Dict, List, Any
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from src.shared.logging import get_logger
|
||||
from src.shared.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class HomeAssistantClient:
|
||||
"""
|
||||
HTTP client for Home Assistant REST API
|
||||
|
||||
Uses long-lived access token authentication via Bearer token.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
token: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize Home Assistant client
|
||||
|
||||
Args:
|
||||
base_url: Home Assistant base URL (default from settings)
|
||||
token: Long-lived access token (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.homeassistant_url).rstrip("/")
|
||||
self.token = token or settings.homeassistant_token
|
||||
self.timeout = timeout
|
||||
|
||||
if not self.token:
|
||||
logger.warning("Home Assistant token not configured")
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with Bearer token authentication"""
|
||||
return {
|
||||
"Authorization": f"Bearer {self.token}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
# ========================================================================
|
||||
# Health & Discovery
|
||||
# ========================================================================
|
||||
|
||||
async def health_check(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Check Home Assistant API connectivity and get version info
|
||||
|
||||
HA Endpoint: GET /api/
|
||||
|
||||
Returns:
|
||||
Dict with connected status, platform name, and version
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
return {
|
||||
"status": "healthy",
|
||||
"connected": True,
|
||||
"platform": "home_assistant",
|
||||
"version": data.get("version", "unknown")
|
||||
}
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"connected": False,
|
||||
"platform": "home_assistant",
|
||||
"error": f"HTTP {response.status_code}"
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Home Assistant health check failed: {e}")
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"connected": False,
|
||||
"platform": "home_assistant",
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
async def get_states(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get all entity states
|
||||
|
||||
HA Endpoint: GET /api/states
|
||||
|
||||
Returns:
|
||||
List of all entity states
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/states",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_state(self, entity_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Get state of a specific entity
|
||||
|
||||
HA Endpoint: GET /api/states/<entity_id>
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., "light.living_room")
|
||||
|
||||
Returns:
|
||||
Entity state dict or None if not found
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/states/{entity_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
if response.status_code == 404:
|
||||
return None
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_config(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Get Home Assistant configuration (includes areas)
|
||||
|
||||
HA Endpoint: GET /api/config
|
||||
|
||||
Returns:
|
||||
Configuration dict including components, location, etc.
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/config",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
# ========================================================================
|
||||
# Device Control
|
||||
# ========================================================================
|
||||
|
||||
async def call_service(
|
||||
self,
|
||||
domain: str,
|
||||
service: str,
|
||||
entity_id: Optional[str] = None,
|
||||
service_data: Optional[Dict[str, Any]] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Call a Home Assistant service
|
||||
|
||||
HA Endpoint: POST /api/services/<domain>/<service>
|
||||
|
||||
Args:
|
||||
domain: Service domain (e.g., "light", "switch", "scene")
|
||||
service: Service name (e.g., "turn_on", "turn_off", "toggle")
|
||||
entity_id: Target entity ID (optional for some services)
|
||||
service_data: Additional service data/attributes
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
payload = service_data.copy() if service_data else {}
|
||||
if entity_id:
|
||||
payload["entity_id"] = entity_id
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/services/{domain}/{service}",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def turn_on(
|
||||
self,
|
||||
entity_id: str,
|
||||
**attributes
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Turn on an entity with optional attributes
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID (e.g., "light.living_room")
|
||||
**attributes: Additional attributes (brightness, color_temp, etc.)
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="turn_on",
|
||||
entity_id=entity_id,
|
||||
service_data=attributes if attributes else None
|
||||
)
|
||||
|
||||
async def turn_off(self, entity_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Turn off an entity
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="turn_off",
|
||||
entity_id=entity_id
|
||||
)
|
||||
|
||||
async def toggle(self, entity_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Toggle an entity
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
domain = entity_id.split(".")[0]
|
||||
return await self.call_service(
|
||||
domain=domain,
|
||||
service="toggle",
|
||||
entity_id=entity_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Scenes
|
||||
# ========================================================================
|
||||
|
||||
async def activate_scene(self, scene_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Activate a scene
|
||||
|
||||
Args:
|
||||
scene_id: Scene entity ID (e.g., "scene.movie_night")
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="scene",
|
||||
service="turn_on",
|
||||
entity_id=scene_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Scripts
|
||||
# ========================================================================
|
||||
|
||||
async def run_script(
|
||||
self,
|
||||
script_id: str,
|
||||
variables: Optional[Dict[str, Any]] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Execute a script with optional variables
|
||||
|
||||
Args:
|
||||
script_id: Script entity ID (e.g., "script.bedtime_routine")
|
||||
variables: Script variables
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
service_data = {"variables": variables} if variables else None
|
||||
return await self.call_service(
|
||||
domain="script",
|
||||
service="turn_on",
|
||||
entity_id=script_id,
|
||||
service_data=service_data
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# Automations
|
||||
# ========================================================================
|
||||
|
||||
async def enable_automation(self, automation_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Enable an automation
|
||||
|
||||
Args:
|
||||
automation_id: Automation entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="automation",
|
||||
service="turn_on",
|
||||
entity_id=automation_id
|
||||
)
|
||||
|
||||
async def disable_automation(self, automation_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Disable an automation
|
||||
|
||||
Args:
|
||||
automation_id: Automation entity ID
|
||||
|
||||
Returns:
|
||||
List of changed states
|
||||
"""
|
||||
return await self.call_service(
|
||||
domain="automation",
|
||||
service="turn_off",
|
||||
entity_id=automation_id
|
||||
)
|
||||
|
||||
# ========================================================================
|
||||
# History
|
||||
# ========================================================================
|
||||
|
||||
async def get_history(
|
||||
self,
|
||||
entity_id: str,
|
||||
hours: int = 24
|
||||
) -> List[List[Dict[str, Any]]]:
|
||||
"""
|
||||
Get state history for an entity
|
||||
|
||||
HA Endpoint: GET /api/history/period/<timestamp>
|
||||
|
||||
Args:
|
||||
entity_id: Entity ID to get history for
|
||||
hours: Number of hours of history (default 24)
|
||||
|
||||
Returns:
|
||||
List of state history entries
|
||||
"""
|
||||
start_time = datetime.now(timezone.utc) - timedelta(hours=hours)
|
||||
timestamp = start_time.isoformat()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/history/period/{timestamp}",
|
||||
headers=self._get_headers(),
|
||||
params={
|
||||
"filter_entity_id": entity_id,
|
||||
"minimal_response": "true"
|
||||
}
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
# ========================================================================
|
||||
# Areas (via template API)
|
||||
# ========================================================================
|
||||
|
||||
async def get_areas(self) -> List[Dict[str, str]]:
|
||||
"""
|
||||
Get all areas/rooms
|
||||
|
||||
Note: The REST API doesn't have a direct areas endpoint.
|
||||
This uses the template API to render area data.
|
||||
|
||||
HA Endpoint: POST /api/template
|
||||
|
||||
Returns:
|
||||
List of area dicts with id and name
|
||||
"""
|
||||
template = """
|
||||
{% set areas_list = [] %}
|
||||
{% for area in areas() %}
|
||||
{% set areas_list = areas_list + [{"id": area, "name": area_name(area)}] %}
|
||||
{% endfor %}
|
||||
{{ areas_list | tojson }}
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/template",
|
||||
headers=self._get_headers(),
|
||||
json={"template": template}
|
||||
)
|
||||
response.raise_for_status()
|
||||
# Response is rendered template as string
|
||||
return json.loads(response.text)
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_homeassistant_client: Optional[HomeAssistantClient] = None
|
||||
|
||||
|
||||
def get_homeassistant_client() -> HomeAssistantClient:
|
||||
"""Get singleton Home Assistant client instance"""
|
||||
global _homeassistant_client
|
||||
if _homeassistant_client is None:
|
||||
_homeassistant_client = HomeAssistantClient()
|
||||
return _homeassistant_client
|
||||
@@ -0,0 +1,383 @@
|
||||
"""
|
||||
Nginx Proxy Manager API Client
|
||||
|
||||
Provides interface to NPM REST API for proxy host and SSL certificate management.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Optional, Dict, List, Any
|
||||
from datetime import datetime, timedelta
|
||||
from src.shared.logging import get_logger
|
||||
from src.shared.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class NPMClient:
|
||||
"""
|
||||
HTTP client for Nginx Proxy Manager API
|
||||
|
||||
Uses JWT Bearer token authentication with automatic token refresh.
|
||||
Tokens expire after ~24 hours.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
email: Optional[str] = None,
|
||||
password: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize NPM client
|
||||
|
||||
Args:
|
||||
base_url: NPM base URL (default from settings)
|
||||
email: NPM admin email (default from settings)
|
||||
password: NPM admin password (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.npm_url).rstrip("/")
|
||||
self.email = email or settings.npm_email
|
||||
self.password = password or settings.npm_password
|
||||
self.timeout = timeout
|
||||
|
||||
self._token: Optional[str] = None
|
||||
self._token_expires: Optional[datetime] = None
|
||||
|
||||
if not self.email or not self.password:
|
||||
logger.warning("NPM credentials not configured")
|
||||
|
||||
async def _ensure_token(self):
|
||||
"""Ensure we have a valid token, refresh if needed"""
|
||||
if self._token and self._token_expires:
|
||||
# If token expires in less than 1 hour, refresh it
|
||||
if datetime.now() + timedelta(hours=1) < self._token_expires:
|
||||
return
|
||||
|
||||
# Get new token
|
||||
await self._refresh_token()
|
||||
|
||||
async def _refresh_token(self):
|
||||
"""Get a new authentication token"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/tokens",
|
||||
json={
|
||||
"identity": self.email,
|
||||
"secret": self.password
|
||||
}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
self._token = data.get("token")
|
||||
# Assume 23-hour expiration to be safe
|
||||
self._token_expires = datetime.now() + timedelta(hours=23)
|
||||
|
||||
logger.info("NPM token refreshed successfully")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to refresh NPM token: {e}")
|
||||
raise
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with authentication"""
|
||||
if not self._token:
|
||||
raise RuntimeError("No NPM token available. Call _ensure_token() first.")
|
||||
|
||||
return {
|
||||
"Authorization": f"Bearer {self._token}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if NPM API is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout, follow_redirects=True) as client:
|
||||
response = await client.get(f"{self.base_url}/api")
|
||||
# Accept any successful response (2xx) or redirect (3xx) as healthy
|
||||
# A redirect indicates the service is up and responding
|
||||
return 200 <= response.status_code < 400
|
||||
except Exception as e:
|
||||
logger.error(f"NPM health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_proxy_hosts(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all proxy hosts
|
||||
|
||||
Returns:
|
||||
List of proxy host configurations
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_proxy_host(self, host_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
Get details of a specific proxy host
|
||||
|
||||
Args:
|
||||
host_id: Proxy host identifier
|
||||
|
||||
Returns:
|
||||
Proxy host configuration
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts/{host_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_proxy_host(
|
||||
self,
|
||||
domain_names: List[str],
|
||||
forward_host: str,
|
||||
forward_port: int,
|
||||
forward_scheme: str = "http",
|
||||
certificate_id: int = 0,
|
||||
ssl_forced: bool = False,
|
||||
block_exploits: bool = True,
|
||||
caching_enabled: bool = True,
|
||||
websocket_upgrade: bool = True,
|
||||
http2_support: bool = True,
|
||||
hsts_enabled: bool = True,
|
||||
advanced_config: str = ""
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a new proxy host
|
||||
|
||||
Args:
|
||||
domain_names: List of domain names for this proxy
|
||||
forward_host: Target host to proxy to
|
||||
forward_port: Target port to proxy to
|
||||
forward_scheme: http or https
|
||||
certificate_id: SSL certificate ID (0 for none)
|
||||
ssl_forced: Force HTTPS redirect
|
||||
block_exploits: Enable exploit blocking
|
||||
caching_enabled: Enable response caching
|
||||
websocket_upgrade: Allow WebSocket upgrades
|
||||
http2_support: Enable HTTP/2
|
||||
hsts_enabled: Enable HSTS headers
|
||||
advanced_config: Custom nginx configuration
|
||||
|
||||
Returns:
|
||||
Created proxy host details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
payload = {
|
||||
"domain_names": domain_names,
|
||||
"forward_scheme": forward_scheme,
|
||||
"forward_host": forward_host,
|
||||
"forward_port": forward_port,
|
||||
"certificate_id": certificate_id,
|
||||
"ssl_forced": ssl_forced,
|
||||
"block_exploits": block_exploits,
|
||||
"caching_enabled": caching_enabled,
|
||||
"allow_websocket_upgrade": websocket_upgrade,
|
||||
"http2_support": http2_support,
|
||||
"hsts_enabled": hsts_enabled,
|
||||
"hsts_subdomains": False,
|
||||
"advanced_config": advanced_config,
|
||||
"access_list_id": 0,
|
||||
"meta": {}
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def update_proxy_host(
|
||||
self,
|
||||
proxy_id: int,
|
||||
config: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update an existing proxy host configuration
|
||||
|
||||
Args:
|
||||
proxy_id: Proxy host ID to update
|
||||
config: Full proxy host configuration (get from get_proxy_host, modify, then update)
|
||||
|
||||
Returns:
|
||||
Updated proxy host details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts/{proxy_id}",
|
||||
headers=self._get_headers(),
|
||||
json=config
|
||||
)
|
||||
|
||||
if not response.is_success:
|
||||
logger.error(f"Update failed: {response.status_code}")
|
||||
logger.error(f"Response: {response.text}")
|
||||
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def enable_authentik_forward_auth(
|
||||
self,
|
||||
proxy_id: int,
|
||||
authentik_url: str = "http://authentik-server:9000"
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Enable Authentik forward authentication on a proxy host
|
||||
|
||||
Args:
|
||||
proxy_id: Proxy host ID to update
|
||||
authentik_url: Authentik server URL (default: http://authentik-server:9000)
|
||||
|
||||
Returns:
|
||||
Updated proxy host details
|
||||
"""
|
||||
# Get current config
|
||||
proxy_host = await self.get_proxy_host(proxy_id)
|
||||
|
||||
# Authentik forward auth configuration
|
||||
auth_config = f"""# Authentik Forward Authentication
|
||||
# Send authentication requests to Authentik
|
||||
auth_request /outpost.goauthentik.io/auth/nginx;
|
||||
|
||||
# Preserve authentication cookies
|
||||
auth_request_set $auth_cookie $upstream_http_set_cookie;
|
||||
add_header Set-Cookie $auth_cookie;
|
||||
|
||||
# Get user information from Authentik
|
||||
auth_request_set $authentik_username $upstream_http_x_authentik_username;
|
||||
auth_request_set $authentik_groups $upstream_http_x_authentik_groups;
|
||||
auth_request_set $authentik_email $upstream_http_x_authentik_email;
|
||||
auth_request_set $authentik_name $upstream_http_x_authentik_name;
|
||||
auth_request_set $authentik_uid $upstream_http_x_authentik_uid;
|
||||
|
||||
# Pass user info to backend
|
||||
proxy_set_header X-authentik-username $authentik_username;
|
||||
proxy_set_header X-authentik-groups $authentik_groups;
|
||||
proxy_set_header X-authentik-email $authentik_email;
|
||||
proxy_set_header X-authentik-name $authentik_name;
|
||||
proxy_set_header X-authentik-uid $authentik_uid;
|
||||
|
||||
# On authentication failure, redirect to Authentik login
|
||||
error_page 401 = @authentik_proxy_signin;
|
||||
|
||||
location @authentik_proxy_signin {{
|
||||
internal;
|
||||
add_header Set-Cookie $auth_cookie;
|
||||
return 302 /outpost.goauthentik.io/start?rd=$scheme://$http_host$request_uri;
|
||||
}}
|
||||
|
||||
# Authentik authentication endpoint
|
||||
location /outpost.goauthentik.io {{
|
||||
proxy_pass {authentik_url}/outpost.goauthentik.io;
|
||||
proxy_set_header X-Original-URL $scheme://$http_host$request_uri;
|
||||
proxy_pass_request_body off;
|
||||
proxy_set_header Content-Length "";
|
||||
proxy_set_header Host $host;
|
||||
}}
|
||||
"""
|
||||
|
||||
# Update the advanced config
|
||||
proxy_host["advanced_config"] = auth_config
|
||||
|
||||
# Remove read-only fields that NPM doesn't accept in updates
|
||||
readonly_fields = [
|
||||
"id", "created_on", "modified_on", "owner", "owner_user_id",
|
||||
"certificate", "use_default_location", "ipv6", "meta", "nginx_online",
|
||||
"nginx_err", "access_list", "certificate_id"
|
||||
]
|
||||
|
||||
clean_config = {k: v for k, v in proxy_host.items() if k not in readonly_fields}
|
||||
|
||||
# Ensure locations is an array (required field)
|
||||
if "locations" not in clean_config or clean_config["locations"] is None:
|
||||
clean_config["locations"] = []
|
||||
|
||||
# Update the proxy host
|
||||
return await self.update_proxy_host(proxy_id, clean_config)
|
||||
|
||||
async def get_certificates(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all SSL certificates
|
||||
|
||||
Returns:
|
||||
List of certificate details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/certificates",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_certificate(
|
||||
self,
|
||||
domain_names: List[str],
|
||||
provider: str = "letsencrypt"
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Request a new SSL certificate from Let's Encrypt
|
||||
|
||||
Args:
|
||||
domain_names: List of domains for the certificate
|
||||
provider: Certificate provider (default: letsencrypt)
|
||||
|
||||
Returns:
|
||||
Certificate details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
payload = {
|
||||
"provider": provider,
|
||||
"domain_names": domain_names,
|
||||
"meta": {
|
||||
"dns_challenge": False
|
||||
}
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/nginx/certificates",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_npm_client: Optional[NPMClient] = None
|
||||
|
||||
|
||||
def get_npm_client() -> NPMClient:
|
||||
"""Get singleton NPM client instance"""
|
||||
global _npm_client
|
||||
if _npm_client is None:
|
||||
_npm_client = NPMClient()
|
||||
return _npm_client
|
||||
@@ -0,0 +1,505 @@
|
||||
"""
|
||||
Portainer API Client
|
||||
|
||||
Provides interface to Portainer REST API for stack and container management.
|
||||
Includes fallback to Docker socket for containers not managed by Portainer.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Optional, Dict, List, Any
|
||||
from src.shared.logging import get_logger
|
||||
from src.shared.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class PortainerClient:
|
||||
"""
|
||||
HTTP client for Portainer API
|
||||
|
||||
Uses access token authentication (X-API-Key header)
|
||||
for long-lived API access without session management.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize Portainer client
|
||||
|
||||
Args:
|
||||
base_url: Portainer base URL (default from settings)
|
||||
api_key: Portainer API access token (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.portainer_url).rstrip("/")
|
||||
self.api_key = api_key or settings.portainer_api_key
|
||||
self.timeout = timeout
|
||||
|
||||
if not self.api_key:
|
||||
logger.warning("Portainer API key not configured")
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with authentication"""
|
||||
return {
|
||||
"X-API-Key": self.api_key,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if Portainer API is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(f"{self.base_url}/api/status")
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"Portainer health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_endpoints(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all Portainer endpoints (Docker environments)
|
||||
|
||||
Returns:
|
||||
List of endpoint configurations
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/endpoints",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_stacks(self, endpoint_id: Optional[int] = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all stacks
|
||||
|
||||
Args:
|
||||
endpoint_id: Filter by specific endpoint (optional)
|
||||
|
||||
Returns:
|
||||
List of stack configurations
|
||||
"""
|
||||
params = {}
|
||||
if endpoint_id:
|
||||
params["endpointId"] = endpoint_id
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks",
|
||||
headers=self._get_headers(),
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_stack(self, stack_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
Get details of a specific stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
|
||||
Returns:
|
||||
Stack configuration details
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_stack(
|
||||
self,
|
||||
name: str,
|
||||
stack_file_content: str,
|
||||
endpoint_id: int
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a new stack from compose file content
|
||||
|
||||
Args:
|
||||
name: Stack name
|
||||
stack_file_content: Docker Compose YAML content
|
||||
endpoint_id: Portainer endpoint to deploy to
|
||||
|
||||
Returns:
|
||||
Created stack details
|
||||
"""
|
||||
payload = {
|
||||
"name": name,
|
||||
"stackFileContent": stack_file_content
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/stacks/create/standalone/string",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def update_stack(
|
||||
self,
|
||||
stack_id: int,
|
||||
stack_file_content: str,
|
||||
endpoint_id: int,
|
||||
prune: bool = False,
|
||||
pull_image: bool = False
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update an existing stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
stack_file_content: New Docker Compose YAML content
|
||||
endpoint_id: Portainer endpoint
|
||||
prune: Remove services no longer defined
|
||||
pull_image: Pull latest images before deployment
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
payload = {
|
||||
"stackFileContent": stack_file_content,
|
||||
"prune": prune,
|
||||
"pullImage": pull_image
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def delete_stack(self, stack_id: int, endpoint_id: int) -> bool:
|
||||
"""
|
||||
Delete a stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.delete(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id}
|
||||
)
|
||||
response.raise_for_status()
|
||||
return True
|
||||
|
||||
async def get_stack_file(self, stack_id: int) -> str:
|
||||
"""
|
||||
Get the compose file content for a stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
|
||||
Returns:
|
||||
Docker Compose YAML content as string
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks/{stack_id}/file",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data.get("StackFileContent", "")
|
||||
|
||||
async def redeploy_stack(
|
||||
self,
|
||||
stack_id: int,
|
||||
endpoint_id: int,
|
||||
pull_image: bool = False
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Redeploy a stack with its current configuration
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
pull_image: Pull latest images before deployment
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
# Get current stack file content
|
||||
stack_content = await self.get_stack_file(stack_id)
|
||||
|
||||
# Get current stack to preserve env vars
|
||||
stack = await self.get_stack(stack_id)
|
||||
env_vars = stack.get("Env", [])
|
||||
|
||||
payload = {
|
||||
"stackFileContent": stack_content,
|
||||
"env": env_vars,
|
||||
"prune": False,
|
||||
"pullImage": pull_image
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def update_stack_env(
|
||||
self,
|
||||
stack_id: int,
|
||||
endpoint_id: int,
|
||||
env_vars: List[Dict[str, str]]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update stack environment variables
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
env_vars: List of {"name": "VAR_NAME", "value": "var_value"} dicts
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
# Get current stack file content (required for update)
|
||||
stack_content = await self.get_stack_file(stack_id)
|
||||
|
||||
payload = {
|
||||
"stackFileContent": stack_content,
|
||||
"env": env_vars,
|
||||
"prune": False,
|
||||
"pullImage": False
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def delete_container(
|
||||
self,
|
||||
endpoint_id: int,
|
||||
container_id: str,
|
||||
force: bool = False
|
||||
) -> bool:
|
||||
"""
|
||||
Delete a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
force: Force remove running container
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
params = {"force": "true" if force else "false"}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.delete(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}",
|
||||
headers=self._get_headers(),
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Deleted container {container_id}")
|
||||
return True
|
||||
|
||||
async def restart_container(self, endpoint_id: int, container_id: str) -> bool:
|
||||
"""
|
||||
Restart a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}/restart",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Restarted container {container_id}")
|
||||
return True
|
||||
|
||||
async def get_containers(self, endpoint_id: int, all_containers: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List containers on a specific endpoint
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
all_containers: Include stopped containers (default: True)
|
||||
|
||||
Returns:
|
||||
List of container details
|
||||
"""
|
||||
params = {"all": 1 if all_containers else 0}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/json",
|
||||
headers=self._get_headers(),
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_container(self, endpoint_id: int, container_id: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get detailed information about a specific container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
|
||||
Returns:
|
||||
Container details including network and port information
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}/json",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def stop_container(self, endpoint_id: int, container_id: str) -> bool:
|
||||
"""
|
||||
Stop a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}/stop",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Stopped container {container_id}")
|
||||
return True
|
||||
|
||||
async def start_container(self, endpoint_id: int, container_id: str) -> bool:
|
||||
"""
|
||||
Start a container
|
||||
|
||||
Args:
|
||||
endpoint_id: Portainer endpoint identifier
|
||||
container_id: Container ID or name
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/endpoints/{endpoint_id}/docker/containers/{container_id}/start",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(f"Started container {container_id}")
|
||||
return True
|
||||
|
||||
# ========================================================================
|
||||
# Helper methods for agent tools (auto-detect endpoint)
|
||||
# ========================================================================
|
||||
|
||||
async def list_containers(self, all_containers: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List containers using auto-detected endpoint
|
||||
|
||||
This is a convenience wrapper that automatically uses the first/default endpoint.
|
||||
|
||||
Args:
|
||||
all_containers: Include stopped containers (default: True)
|
||||
|
||||
Returns:
|
||||
List of container details
|
||||
"""
|
||||
endpoints = await self.get_endpoints()
|
||||
if not endpoints:
|
||||
raise RuntimeError("No Portainer endpoints available")
|
||||
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
return await self.get_containers(endpoint_id, all_containers)
|
||||
|
||||
async def inspect_container(self, container_name: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Inspect a container by name using auto-detected endpoint
|
||||
|
||||
This is a convenience wrapper that automatically uses the first/default endpoint.
|
||||
|
||||
Args:
|
||||
container_name: Container name (e.g., "jellyfin", "ollama")
|
||||
|
||||
Returns:
|
||||
Container details or None if not found
|
||||
"""
|
||||
endpoints = await self.get_endpoints()
|
||||
if not endpoints:
|
||||
raise RuntimeError("No Portainer endpoints available")
|
||||
|
||||
endpoint_id = endpoints[0]["Id"]
|
||||
|
||||
# List all containers to find the one matching the name
|
||||
all_containers = await self.get_containers(endpoint_id, all_containers=True)
|
||||
|
||||
for container in all_containers:
|
||||
# Container names come as array like ['/jellyfin']
|
||||
names = container.get('Names', [])
|
||||
for name in names:
|
||||
clean_name = name.lstrip('/')
|
||||
if clean_name == container_name or clean_name.lower() == container_name.lower():
|
||||
# Get detailed info using container ID
|
||||
container_id = container['Id']
|
||||
return await self.get_container(endpoint_id, container_id)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_portainer_client: Optional[PortainerClient] = None
|
||||
|
||||
|
||||
def get_portainer_client() -> PortainerClient:
|
||||
"""Get singleton Portainer client instance"""
|
||||
global _portainer_client
|
||||
if _portainer_client is None:
|
||||
_portainer_client = PortainerClient()
|
||||
return _portainer_client
|
||||
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
Global configuration for Core Code API
|
||||
|
||||
All configuration is loaded from environment variables or .env file.
|
||||
See .env.example for available settings.
|
||||
"""
|
||||
import tomllib
|
||||
from pathlib import Path
|
||||
from pydantic_settings import BaseSettings
|
||||
from functools import lru_cache
|
||||
|
||||
|
||||
def _get_version_from_pyproject() -> str:
|
||||
"""Load version from pyproject.toml"""
|
||||
pyproject_path = Path(__file__).parent.parent.parent / "pyproject.toml"
|
||||
try:
|
||||
with open(pyproject_path, "rb") as f:
|
||||
data = tomllib.load(f)
|
||||
return data.get("project", {}).get("version", "0.0.0")
|
||||
except FileNotFoundError:
|
||||
return "0.0.0"
|
||||
|
||||
|
||||
__version__ = _get_version_from_pyproject()
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Global application settings"""
|
||||
|
||||
# Application
|
||||
app_name: str = "Core Code API"
|
||||
app_version: str = __version__
|
||||
debug: bool = False
|
||||
|
||||
# Server
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8083
|
||||
|
||||
# CORS
|
||||
cors_origins: list[str] = ["*"]
|
||||
cors_credentials: bool = True
|
||||
cors_methods: list[str] = ["*"]
|
||||
cors_headers: list[str] = ["*"]
|
||||
|
||||
# Logging
|
||||
log_level: str = "DEBUG"
|
||||
|
||||
# Ollama Configuration (for AI orchestration)
|
||||
ollama_base_url: str # Required - set OLLAMA_BASE_URL in .env
|
||||
ollama_timeout: int = 300 # 5 minutes
|
||||
|
||||
# Model Configuration
|
||||
default_model: str = "mistral-nemo-large:latest"
|
||||
agent_model: str = "mistral-nemo-large:latest"
|
||||
code_models: str = "mistral-nemo-large:latest"
|
||||
|
||||
# System Prompt Variant (for A/B testing)
|
||||
system_prompt_variant: str = "v8_holistic"
|
||||
|
||||
# Agent Configuration
|
||||
agent_fallback_enabled: bool = True
|
||||
|
||||
# Model Aliases (OpenAI → Local)
|
||||
alias_gpt35: str = "gemma:7b"
|
||||
alias_gpt4: str = "mistral:7b"
|
||||
alias_gpt4_turbo: str = "mixtral:8x7b"
|
||||
alias_gpt4_code: str = "codestral:latest"
|
||||
|
||||
# Memory Configuration
|
||||
memory_tier1_max_turns: int = 10
|
||||
memory_consolidation_threshold: int = 10
|
||||
|
||||
# Qdrant Configuration
|
||||
qdrant_host: str = "qdrant"
|
||||
qdrant_port: int = 6333
|
||||
qdrant_collection_conversations: str = "core_api_conversations"
|
||||
qdrant_collection_documents: str = "core_api_documents"
|
||||
qdrant_collection_user_facts: str = "core_api_user_facts"
|
||||
|
||||
# Embeddings (using Ollama)
|
||||
embedding_model: str = "nomic-embed-text"
|
||||
embedding_dimension: int = 768
|
||||
embedding_batch_size: int = 32
|
||||
|
||||
# Search Configuration
|
||||
search_provider: str = "searxng"
|
||||
searxng_url: str # Required - set SEARXNG_URL in .env
|
||||
|
||||
# Infrastructure Management (Portainer)
|
||||
portainer_url: str # Required
|
||||
portainer_api_key: str # Required
|
||||
|
||||
# Infrastructure Management (Nginx Proxy Manager)
|
||||
npm_url: str # Required
|
||||
npm_email: str # Required
|
||||
npm_password: str # Required
|
||||
|
||||
# Home Assistant Configuration
|
||||
homeassistant_url: str # Required
|
||||
homeassistant_token: str # Required
|
||||
homeassistant_timeout: int = 30
|
||||
|
||||
# PostgreSQL Database
|
||||
postgres_host: str # Required
|
||||
postgres_user: str = "core_api"
|
||||
postgres_password: str # Required
|
||||
postgres_database: str = "core_api"
|
||||
|
||||
@property
|
||||
def database_url(self) -> str:
|
||||
"""Construct database URL from components"""
|
||||
return f"postgresql://{self.postgres_user}:{self.postgres_password}@{self.postgres_host}/{self.postgres_database}"
|
||||
|
||||
# OIDC Authentication (Authentik)
|
||||
oidc_enabled: bool = False
|
||||
oidc_issuer: str = "https://auth.schweitz.net/application/o/core-api/"
|
||||
oidc_audience: str = "core-api"
|
||||
|
||||
# Authentik API (for token validation and user management)
|
||||
authentik_url: str = "https://auth.schweitz.net"
|
||||
authentik_username: str = ""
|
||||
authentik_password: str = ""
|
||||
|
||||
@property
|
||||
def model_aliases(self) -> dict:
|
||||
"""Computed property for model aliases"""
|
||||
return {
|
||||
"gpt-3.5-turbo": self.alias_gpt35,
|
||||
"gpt-4": self.alias_gpt4,
|
||||
"gpt-4-turbo": self.alias_gpt4_turbo,
|
||||
"gpt-4-code": self.alias_gpt4_code,
|
||||
}
|
||||
|
||||
def get_code_models(self) -> list[str]:
|
||||
"""Parse comma-separated code models"""
|
||||
return [m.strip().strip('"').strip("'") for m in self.code_models.split(",") if m.strip()]
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = False
|
||||
extra = "ignore"
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_settings() -> Settings:
|
||||
"""Cached settings instance"""
|
||||
return Settings()
|
||||
@@ -0,0 +1,135 @@
|
||||
"""
|
||||
Database Connection Module
|
||||
|
||||
Provides async PostgreSQL connectivity using SQLAlchemy 2.0 with asyncpg driver.
|
||||
"""
|
||||
from typing import AsyncGenerator, Optional
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncSession,
|
||||
AsyncEngine,
|
||||
create_async_engine,
|
||||
async_sessionmaker,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from src.shared.config import get_settings
|
||||
from src.shared.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
"""
|
||||
SQLAlchemy declarative base for all models
|
||||
|
||||
All database models should inherit from this class.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class Database:
|
||||
"""
|
||||
Async database connection manager
|
||||
|
||||
Provides async engine and session factory for PostgreSQL connections.
|
||||
"""
|
||||
|
||||
def __init__(self, database_url: Optional[str] = None):
|
||||
"""Initialize database connection manager"""
|
||||
url = database_url or settings.database_url
|
||||
if url.startswith("postgresql://"):
|
||||
url = url.replace("postgresql://", "postgresql+asyncpg://", 1)
|
||||
|
||||
self._url = url
|
||||
self._engine: Optional[AsyncEngine] = None
|
||||
self._session_factory: Optional[async_sessionmaker[AsyncSession]] = None
|
||||
|
||||
@property
|
||||
def engine(self) -> AsyncEngine:
|
||||
"""Get or create the async database engine"""
|
||||
if self._engine is None:
|
||||
self._engine = create_async_engine(
|
||||
self._url,
|
||||
echo=settings.debug,
|
||||
poolclass=NullPool,
|
||||
)
|
||||
logger.info(f"Database engine created for {self._url.split('@')[-1]}")
|
||||
return self._engine
|
||||
|
||||
@property
|
||||
def session_factory(self) -> async_sessionmaker[AsyncSession]:
|
||||
"""Get or create the async session factory"""
|
||||
if self._session_factory is None:
|
||||
self._session_factory = async_sessionmaker(
|
||||
bind=self.engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
autocommit=False,
|
||||
autoflush=False,
|
||||
)
|
||||
return self._session_factory
|
||||
|
||||
async def create_tables(self) -> None:
|
||||
"""Create all database tables (dev/testing only)"""
|
||||
async with self.engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("Database tables created")
|
||||
|
||||
async def drop_tables(self) -> None:
|
||||
"""Drop all database tables (WARNING: destroys data)"""
|
||||
async with self.engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
logger.warning("Database tables dropped")
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""Check if database connection is healthy"""
|
||||
try:
|
||||
async with self.session_factory() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Database health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close database connections"""
|
||||
if self._engine is not None:
|
||||
await self._engine.dispose()
|
||||
self._engine = None
|
||||
self._session_factory = None
|
||||
logger.info("Database connections closed")
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_database: Optional[Database] = None
|
||||
|
||||
|
||||
def get_database() -> Database:
|
||||
"""Get singleton database instance"""
|
||||
global _database
|
||||
if _database is None:
|
||||
_database = Database()
|
||||
return _database
|
||||
|
||||
|
||||
async def get_async_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""
|
||||
FastAPI dependency for database sessions
|
||||
|
||||
Usage:
|
||||
@router.get("/items")
|
||||
async def get_items(session: AsyncSession = Depends(get_async_session)):
|
||||
result = await session.execute(select(Item))
|
||||
return result.scalars().all()
|
||||
"""
|
||||
database = get_database()
|
||||
async with database.session_factory() as session:
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
@@ -0,0 +1,37 @@
|
||||
"""
|
||||
Logging configuration for Core Code API
|
||||
"""
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def setup_logging(log_level: str = "INFO") -> None:
|
||||
"""
|
||||
Configure logging for the application
|
||||
|
||||
Args:
|
||||
log_level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL)
|
||||
"""
|
||||
log_dir = Path("logs")
|
||||
log_dir.mkdir(exist_ok=True)
|
||||
|
||||
logging.basicConfig(
|
||||
level=getattr(logging, log_level.upper()),
|
||||
format="%(asctime)s | %(levelname)-8s | %(name)s:%(funcName)s:%(lineno)d | %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
handlers=[
|
||||
logging.StreamHandler(sys.stdout),
|
||||
logging.FileHandler(log_dir / "app.log", encoding="utf-8")
|
||||
]
|
||||
)
|
||||
|
||||
# Set specific log levels for third-party libraries
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def get_logger(name: str) -> logging.Logger:
|
||||
"""Get a logger instance"""
|
||||
return logging.getLogger(name)
|
||||
@@ -0,0 +1,31 @@
|
||||
"""
|
||||
Security initialization module
|
||||
|
||||
Handles OIDC configuration and authentication setup
|
||||
"""
|
||||
from src.shared.config import Settings
|
||||
from src.shared.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def initialize_oidc(settings: Settings) -> None:
|
||||
"""
|
||||
Initialize OIDC authentication configuration
|
||||
|
||||
Args:
|
||||
settings: Application settings containing OIDC configuration
|
||||
"""
|
||||
# Import here to avoid circular imports
|
||||
from src.auth.oidc import oidc_config
|
||||
|
||||
oidc_config.configure(
|
||||
enabled=settings.oidc_enabled,
|
||||
issuer=settings.oidc_issuer,
|
||||
audience=settings.oidc_audience
|
||||
)
|
||||
|
||||
if settings.oidc_enabled:
|
||||
logger.info(f"OIDC authentication enabled (issuer: {settings.oidc_issuer})")
|
||||
else:
|
||||
logger.info("OIDC authentication disabled - API is publicly accessible")
|
||||
@@ -1,13 +0,0 @@
|
||||
"""
|
||||
Web scraper module for extracting content from websites
|
||||
"""
|
||||
from src.web_scraper.router import router
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.service import WebScraperService
|
||||
|
||||
__all__ = [
|
||||
"router",
|
||||
"WebScraperRequest",
|
||||
"WebScraperResponse",
|
||||
"WebScraperService",
|
||||
]
|
||||
@@ -1,32 +0,0 @@
|
||||
"""
|
||||
Configuration for web scraper module
|
||||
"""
|
||||
from pydantic_settings import BaseSettings
|
||||
from functools import lru_cache
|
||||
|
||||
|
||||
class WebScraperSettings(BaseSettings):
|
||||
"""Web scraper specific settings"""
|
||||
|
||||
# HTTP client configuration
|
||||
request_timeout: int = 30
|
||||
max_redirects: int = 5
|
||||
user_agent: str = "Mozilla/5.0 (compatible; CoreCode/1.0)"
|
||||
|
||||
# Content extraction
|
||||
default_max_length: int = 10000
|
||||
max_links_to_extract: int = 50
|
||||
|
||||
# Rate limiting (future use)
|
||||
rate_limit_enabled: bool = False
|
||||
requests_per_minute: int = 60
|
||||
|
||||
class Config:
|
||||
env_prefix = "WEB_SCRAPER_"
|
||||
case_sensitive = False
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_web_scraper_settings() -> WebScraperSettings:
|
||||
"""Cached web scraper settings instance"""
|
||||
return WebScraperSettings()
|
||||
@@ -1,18 +0,0 @@
|
||||
"""
|
||||
Custom exceptions for web scraper module
|
||||
"""
|
||||
|
||||
|
||||
class WebScraperException(Exception):
|
||||
"""Base exception for web scraper module"""
|
||||
pass
|
||||
|
||||
|
||||
class FetchError(WebScraperException):
|
||||
"""Raised when URL fetch fails"""
|
||||
pass
|
||||
|
||||
|
||||
class ScrapingError(WebScraperException):
|
||||
"""Raised when content extraction fails"""
|
||||
pass
|
||||
@@ -1,79 +0,0 @@
|
||||
"""
|
||||
API routes for web scraper module
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from src.logging_config import get_logger
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.service import WebScraperService
|
||||
from src.web_scraper.exceptions import FetchError, ScrapingError
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
router = APIRouter(
|
||||
prefix="/web-scraper",
|
||||
tags=["Web Scraper"]
|
||||
)
|
||||
|
||||
# Initialize service (could be dependency injected for testing)
|
||||
scraper_service = WebScraperService()
|
||||
|
||||
|
||||
@router.post(
|
||||
"/scrape",
|
||||
response_model=WebScraperResponse,
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Scrape website content",
|
||||
description="""
|
||||
Scrape and extract main content from a website.
|
||||
|
||||
Uses trafilatura for intelligent content extraction (articles, blog posts, documentation),
|
||||
with BeautifulSoup as fallback. Perfect for feeding webpage content to LLMs.
|
||||
|
||||
**Features:**
|
||||
- Intelligent main content extraction
|
||||
- Removes navigation, ads, footers
|
||||
- Optional link extraction
|
||||
- Configurable content length limits
|
||||
|
||||
**Rate Limiting:** None (internal network use only)
|
||||
"""
|
||||
)
|
||||
async def scrape_website(request: WebScraperRequest) -> WebScraperResponse:
|
||||
"""
|
||||
Scrape a website and extract its main content
|
||||
|
||||
Args:
|
||||
request: Scraping request with URL and options
|
||||
|
||||
Returns:
|
||||
Extracted content with metadata
|
||||
|
||||
Raises:
|
||||
HTTPException: 400 for fetch errors, 500 for processing errors
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Received scrape request for: {request.url}")
|
||||
result = await scraper_service.scrape_url(request)
|
||||
return result
|
||||
|
||||
except FetchError as e:
|
||||
logger.warning(f"Fetch failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Failed to fetch URL: {str(e)}"
|
||||
)
|
||||
|
||||
except ScrapingError as e:
|
||||
logger.error(f"Scraping failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to extract content: {str(e)}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error: {str(e)}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="An unexpected error occurred"
|
||||
)
|
||||
@@ -1,69 +0,0 @@
|
||||
"""
|
||||
Pydantic schemas for web scraper module
|
||||
"""
|
||||
from pydantic import HttpUrl, Field
|
||||
from typing import Optional
|
||||
from datetime import datetime
|
||||
from src.base_schema import BaseSchema
|
||||
|
||||
|
||||
class WebScraperRequest(BaseSchema):
|
||||
"""Request model for web scraping"""
|
||||
|
||||
url: HttpUrl = Field(
|
||||
...,
|
||||
description="The URL to scrape",
|
||||
examples=["https://example.com/article"]
|
||||
)
|
||||
|
||||
extract_main_content: bool = Field(
|
||||
default=True,
|
||||
description="Use intelligent content extraction (trafilatura) vs raw HTML parsing"
|
||||
)
|
||||
|
||||
include_links: bool = Field(
|
||||
default=False,
|
||||
description="Include list of links found on the page"
|
||||
)
|
||||
|
||||
max_length: Optional[int] = Field(
|
||||
default=10000,
|
||||
ge=100,
|
||||
le=100000,
|
||||
description="Maximum content length to return (100-100000 chars)"
|
||||
)
|
||||
|
||||
|
||||
class WebScraperResponse(BaseSchema):
|
||||
"""Response model for web scraping"""
|
||||
|
||||
url: str = Field(
|
||||
...,
|
||||
description="The scraped URL"
|
||||
)
|
||||
|
||||
title: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Page title extracted from <title> tag"
|
||||
)
|
||||
|
||||
content: str = Field(
|
||||
...,
|
||||
description="Extracted page content"
|
||||
)
|
||||
|
||||
extracted_at: datetime = Field(
|
||||
...,
|
||||
description="UTC timestamp when content was extracted"
|
||||
)
|
||||
|
||||
content_length: int = Field(
|
||||
...,
|
||||
ge=0,
|
||||
description="Length of extracted content in characters"
|
||||
)
|
||||
|
||||
links: Optional[list[str]] = Field(
|
||||
default=None,
|
||||
description="List of HTTP(S) links found on the page (max 50)"
|
||||
)
|
||||
@@ -1,213 +0,0 @@
|
||||
"""
|
||||
Business logic for web scraper module
|
||||
"""
|
||||
import httpx
|
||||
from bs4 import BeautifulSoup
|
||||
import trafilatura
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from src.logging_config import get_logger
|
||||
from src.web_scraper.config import get_web_scraper_settings
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.exceptions import ScrapingError, FetchError
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class WebScraperService:
|
||||
"""Service class for web scraping operations"""
|
||||
|
||||
def __init__(self):
|
||||
self.settings = get_web_scraper_settings()
|
||||
|
||||
async def scrape_url(self, request: WebScraperRequest) -> WebScraperResponse:
|
||||
"""
|
||||
Scrape and extract content from a URL
|
||||
|
||||
Args:
|
||||
request: Scraping request parameters
|
||||
|
||||
Returns:
|
||||
Extracted content with metadata
|
||||
|
||||
Raises:
|
||||
FetchError: If URL cannot be fetched
|
||||
ScrapingError: If content extraction fails
|
||||
"""
|
||||
url_str = str(request.url)
|
||||
logger.info(f"Starting scrape for URL: {url_str}")
|
||||
|
||||
try:
|
||||
# Fetch the webpage
|
||||
html_content = await self._fetch_url(url_str)
|
||||
|
||||
# Extract content based on settings
|
||||
if request.extract_main_content:
|
||||
content = self._extract_main_content(html_content, request.include_links)
|
||||
else:
|
||||
content = self._extract_basic_content(html_content)
|
||||
|
||||
# Extract metadata
|
||||
title = self._extract_title(html_content)
|
||||
links = self._extract_links(html_content) if request.include_links else None
|
||||
|
||||
# Clean and truncate content
|
||||
content = self._clean_content(content)
|
||||
if request.max_length and len(content) > request.max_length:
|
||||
content = content[:request.max_length] + "\n\n[Content truncated...]"
|
||||
logger.debug(f"Content truncated to {request.max_length} characters")
|
||||
|
||||
logger.info(f"Successfully scraped {len(content)} characters from {url_str}")
|
||||
|
||||
return WebScraperResponse(
|
||||
url=url_str,
|
||||
title=title,
|
||||
content=content,
|
||||
extracted_at=datetime.now(timezone.utc),
|
||||
content_length=len(content),
|
||||
links=links
|
||||
)
|
||||
|
||||
except FetchError:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Scraping failed for {url_str}: {str(e)}", exc_info=True)
|
||||
raise ScrapingError(f"Failed to scrape content: {str(e)}")
|
||||
|
||||
async def _fetch_url(self, url: str) -> str:
|
||||
"""
|
||||
Fetch HTML content from URL
|
||||
|
||||
Args:
|
||||
url: URL to fetch
|
||||
|
||||
Returns:
|
||||
HTML content as string
|
||||
|
||||
Raises:
|
||||
FetchError: If fetch fails
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.settings.request_timeout,
|
||||
follow_redirects=True,
|
||||
max_redirects=self.settings.max_redirects
|
||||
) as client:
|
||||
logger.debug(f"Fetching URL: {url}")
|
||||
response = await client.get(
|
||||
url,
|
||||
headers={"User-Agent": self.settings.user_agent}
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.debug(f"Fetched {len(response.text)} bytes from {url}")
|
||||
return response.text
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
logger.error(f"HTTP error {e.response.status_code} for {url}")
|
||||
raise FetchError(f"HTTP {e.response.status_code}: {e.response.reason_phrase}")
|
||||
except httpx.RequestError as e:
|
||||
logger.error(f"Request error for {url}: {str(e)}")
|
||||
raise FetchError(f"Failed to fetch URL: {str(e)}")
|
||||
|
||||
def _extract_main_content(self, html: str, include_links: bool = False) -> str:
|
||||
"""
|
||||
Extract main content using trafilatura (intelligent extraction)
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
include_links: Whether to preserve links in output
|
||||
|
||||
Returns:
|
||||
Extracted content
|
||||
"""
|
||||
logger.debug("Extracting main content with trafilatura")
|
||||
content = trafilatura.extract(
|
||||
html,
|
||||
include_links=include_links,
|
||||
include_images=False,
|
||||
output_format='txt',
|
||||
no_fallback=False
|
||||
)
|
||||
|
||||
# Fallback to BeautifulSoup if trafilatura fails
|
||||
if not content:
|
||||
logger.debug("Trafilatura extraction failed, falling back to BeautifulSoup")
|
||||
content = self._extract_basic_content(html)
|
||||
|
||||
return content
|
||||
|
||||
def _extract_basic_content(self, html: str) -> str:
|
||||
"""
|
||||
Extract content using basic BeautifulSoup parsing
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
Extracted text content
|
||||
"""
|
||||
logger.debug("Extracting content with BeautifulSoup")
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
|
||||
# Remove unwanted elements
|
||||
for element in soup(["script", "style", "nav", "footer", "header", "aside"]):
|
||||
element.decompose()
|
||||
|
||||
# Extract text
|
||||
text = soup.get_text(separator='\n', strip=True)
|
||||
return text
|
||||
|
||||
def _extract_title(self, html: str) -> Optional[str]:
|
||||
"""
|
||||
Extract page title from HTML
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
Page title or None
|
||||
"""
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
title = soup.title.string if soup.title else None
|
||||
if title:
|
||||
title = title.strip()
|
||||
logger.debug(f"Extracted title: {title}")
|
||||
return title
|
||||
|
||||
def _extract_links(self, html: str) -> list[str]:
|
||||
"""
|
||||
Extract HTTP(S) links from HTML
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
List of absolute HTTP(S) URLs
|
||||
"""
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
links = [
|
||||
a.get('href')
|
||||
for a in soup.find_all('a', href=True)
|
||||
if a.get('href', '').startswith('http')
|
||||
]
|
||||
|
||||
# Limit number of links
|
||||
links = links[:self.settings.max_links_to_extract]
|
||||
logger.debug(f"Extracted {len(links)} links")
|
||||
return links
|
||||
|
||||
def _clean_content(self, content: str) -> str:
|
||||
"""
|
||||
Clean and normalize extracted content
|
||||
|
||||
Args:
|
||||
content: Raw extracted content
|
||||
|
||||
Returns:
|
||||
Cleaned content
|
||||
"""
|
||||
# Remove empty lines and normalize whitespace
|
||||
lines = [line.strip() for line in content.split('\n') if line.strip()]
|
||||
cleaned = '\n'.join(lines)
|
||||
return cleaned
|
||||
@@ -71,14 +71,14 @@
|
||||
font-weight: 600;
|
||||
color: #fff;
|
||||
text-transform: capitalize;
|
||||
min-width: 300px;
|
||||
min-width: 200px;
|
||||
}
|
||||
|
||||
.service-status {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
min-width: 120px;
|
||||
min-width: 180px;
|
||||
}
|
||||
|
||||
.status-indicator {
|
||||
@@ -103,54 +103,6 @@
|
||||
color: #a0a0a0;
|
||||
}
|
||||
|
||||
.uptime-status {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
min-width: 150px;
|
||||
padding: 4px 10px;
|
||||
background: rgba(0, 0, 0, 0.2);
|
||||
border-radius: 4px;
|
||||
cursor: pointer;
|
||||
transition: background 0.2s;
|
||||
text-decoration: none;
|
||||
color: inherit;
|
||||
}
|
||||
|
||||
.uptime-status:hover {
|
||||
background: rgba(0, 0, 0, 0.4);
|
||||
}
|
||||
|
||||
.uptime-percentage {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.uptime-percentage.excellent {
|
||||
color: #48bb78;
|
||||
}
|
||||
|
||||
.uptime-percentage.good {
|
||||
color: #68d391;
|
||||
}
|
||||
|
||||
.uptime-percentage.warning {
|
||||
color: #ed8936;
|
||||
}
|
||||
|
||||
.uptime-percentage.critical {
|
||||
color: #f56565;
|
||||
}
|
||||
|
||||
.uptime-percentage.unknown {
|
||||
color: #718096;
|
||||
}
|
||||
|
||||
.uptime-icon {
|
||||
font-size: 11px;
|
||||
color: #a0a0a0;
|
||||
}
|
||||
|
||||
.service-right {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
@@ -248,10 +200,6 @@
|
||||
min-width: 100px;
|
||||
}
|
||||
|
||||
.uptime-status {
|
||||
min-width: 100px;
|
||||
}
|
||||
|
||||
.service-right {
|
||||
width: 100%;
|
||||
justify-content: flex-end;
|
||||
@@ -264,12 +212,12 @@
|
||||
<div id="error-container"></div>
|
||||
|
||||
<div class="section">
|
||||
<div class="section-header">🎛️ Stoppable Services</div>
|
||||
<div class="section-header">Stoppable Services</div>
|
||||
<div id="stoppable-container" class="loading">Loading services...</div>
|
||||
</div>
|
||||
|
||||
<div class="section">
|
||||
<div class="section-header">🔒 Always-On Infrastructure</div>
|
||||
<div class="section-header">Always-On Infrastructure</div>
|
||||
<div id="always-on-container" class="loading">Loading infrastructure...</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -277,11 +225,9 @@
|
||||
<script>
|
||||
// Use relative URL to work in any context (iframe, direct access, etc.)
|
||||
const API_BASE = '';
|
||||
const KUMA_BASE = window.location.protocol + '//' + window.location.hostname + ':3001';
|
||||
|
||||
let services = [];
|
||||
let alwaysOnServices = [];
|
||||
let monitors = {};
|
||||
|
||||
async function fetchData() {
|
||||
try {
|
||||
@@ -306,27 +252,13 @@
|
||||
alwaysOnServices = data.service_groups.always_on;
|
||||
}
|
||||
|
||||
// Build monitors map
|
||||
const monitorsMap = {};
|
||||
if (data.monitors) {
|
||||
data.monitors.forEach(monitor => {
|
||||
const name = monitor.name.toLowerCase().replace(/[^a-z0-9]/g, '-');
|
||||
monitorsMap[name] = {
|
||||
id: monitor.id,
|
||||
uptime_24h: monitor.uptime_24h || 0,
|
||||
active: monitor.active !== false
|
||||
};
|
||||
});
|
||||
}
|
||||
monitors = monitorsMap;
|
||||
|
||||
renderServices();
|
||||
document.getElementById('error-container').innerHTML = '';
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error fetching data:', error);
|
||||
document.getElementById('error-container').innerHTML =
|
||||
`<div class="error">❌ Failed to connect to API: ${error.message}</div>`;
|
||||
`<div class="error">Failed to connect to API: ${error.message}</div>`;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -334,43 +266,9 @@
|
||||
return alwaysOnServices.includes(serviceName.toLowerCase());
|
||||
}
|
||||
|
||||
function getUptimeInfo(serviceName) {
|
||||
const monitorKey = serviceName.toLowerCase().replace(/[^a-z0-9]/g, '-');
|
||||
const monitor = monitors[monitorKey];
|
||||
|
||||
if (!monitor) {
|
||||
return {
|
||||
percentage: 0,
|
||||
class: 'unknown',
|
||||
text: 'No monitor',
|
||||
id: null
|
||||
};
|
||||
}
|
||||
|
||||
const uptime = monitor.uptime_24h;
|
||||
let className = 'unknown';
|
||||
|
||||
if (uptime >= 99.5) className = 'excellent';
|
||||
else if (uptime >= 95) className = 'good';
|
||||
else if (uptime >= 90) className = 'warning';
|
||||
else if (uptime > 0) className = 'critical';
|
||||
|
||||
return {
|
||||
percentage: uptime,
|
||||
class: className,
|
||||
text: uptime > 0 ? `${uptime.toFixed(1)}% ↑` : 'Down',
|
||||
id: monitor.id
|
||||
};
|
||||
}
|
||||
|
||||
function renderServiceRow(service) {
|
||||
const isRunning = service.containers_running > 0;
|
||||
const alwaysOn = isAlwaysOn(service.name);
|
||||
const uptime = getUptimeInfo(service.name);
|
||||
|
||||
const kumaLink = uptime.id ?
|
||||
`${KUMA_BASE}/dashboard/${uptime.id}` :
|
||||
KUMA_BASE;
|
||||
|
||||
return `
|
||||
<div class="service-row" data-service="${service.name}">
|
||||
@@ -385,10 +283,6 @@
|
||||
${isRunning ? `Running (${service.containers_running}/${service.containers_total})` : 'Stopped'}
|
||||
</span>
|
||||
</div>
|
||||
<a href="${kumaLink}" target="_blank" class="uptime-status" title="View in Uptime Kuma">
|
||||
<span class="uptime-icon">📊</span>
|
||||
<span class="uptime-percentage ${uptime.class}">${uptime.text}</span>
|
||||
</a>
|
||||
</div>
|
||||
<div class="service-right">
|
||||
<button class="btn btn-start"
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# Tests package
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,627 @@
|
||||
"""Tests for authentication service."""
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.domains.auth.models import User, Role, Group, UserPreferences
|
||||
from src.domains.auth.schemas import TokenInfoSchema, RoleSchema
|
||||
from src.domains.auth.service import AuthService, get_auth_service
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Fixtures
|
||||
# =============================================================================
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session():
|
||||
"""Create a mock async database session."""
|
||||
session = AsyncMock(spec=AsyncSession)
|
||||
session.execute = AsyncMock()
|
||||
session.commit = AsyncMock()
|
||||
session.flush = AsyncMock()
|
||||
session.refresh = AsyncMock()
|
||||
session.add = MagicMock()
|
||||
return session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def auth_service(mock_session):
|
||||
"""Create an AuthService instance with mock session."""
|
||||
return AuthService(mock_session)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_user():
|
||||
"""Create a sample user for testing."""
|
||||
user = User(
|
||||
id=uuid.uuid4(),
|
||||
authentik_id=uuid.uuid4(),
|
||||
email="test@example.com",
|
||||
name="Test User",
|
||||
avatar_url="https://example.com/avatar.jpg",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
last_login=datetime.now(timezone.utc),
|
||||
)
|
||||
user.roles = []
|
||||
user.preferences = None
|
||||
return user
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_role():
|
||||
"""Create a sample role for testing."""
|
||||
return Role(
|
||||
id=uuid.uuid4(),
|
||||
name="control-room.general:admin",
|
||||
domain="control-room",
|
||||
category="general",
|
||||
action="admin",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_group(sample_role):
|
||||
"""Create a sample group for testing."""
|
||||
group = Group(
|
||||
id=uuid.uuid4(),
|
||||
authentik_id=uuid.uuid4(),
|
||||
name="Administrators",
|
||||
is_superuser=True,
|
||||
parent_name=None,
|
||||
member_count=5,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
synced_at=datetime.now(timezone.utc),
|
||||
)
|
||||
group.roles = [sample_role]
|
||||
return group
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_token_info():
|
||||
"""Create sample token info from Authentik."""
|
||||
return TokenInfoSchema(
|
||||
sub=str(uuid.uuid4()),
|
||||
email="test@example.com",
|
||||
name="Test User",
|
||||
preferred_username="testuser",
|
||||
groups=["Administrators", "Developers"],
|
||||
picture="https://example.com/avatar.jpg",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_preferences():
|
||||
"""Create sample user preferences."""
|
||||
return UserPreferences(
|
||||
user_id=uuid.uuid4(),
|
||||
theme="dark",
|
||||
default_room="control-room",
|
||||
preferences_json={"notifications": True},
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# AuthService Initialization Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestAuthServiceInit:
|
||||
"""Test AuthService initialization."""
|
||||
|
||||
def test_init_with_session(self, mock_session):
|
||||
"""AuthService should initialize with session."""
|
||||
service = AuthService(mock_session)
|
||||
assert service.session is mock_session
|
||||
|
||||
def test_init_sets_userinfo_url(self, mock_session):
|
||||
"""AuthService should set userinfo URL from settings."""
|
||||
service = AuthService(mock_session)
|
||||
assert "userinfo" in service.userinfo_url
|
||||
|
||||
def test_get_auth_service_factory(self, mock_session):
|
||||
"""get_auth_service should return AuthService instance."""
|
||||
service = get_auth_service(mock_session)
|
||||
assert isinstance(service, AuthService)
|
||||
assert service.session is mock_session
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Token Validation Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestValidateToken:
|
||||
"""Test token validation via Authentik userinfo endpoint."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_token_success(self, auth_service):
|
||||
"""validate_token should return token info on success."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"sub": str(uuid.uuid4()),
|
||||
"email": "test@example.com",
|
||||
"name": "Test User",
|
||||
"preferred_username": "testuser",
|
||||
"groups": ["Administrators"],
|
||||
"picture": "https://example.com/avatar.jpg",
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
|
||||
result = await auth_service.validate_token("valid_token")
|
||||
|
||||
assert result.email == "test@example.com"
|
||||
assert result.name == "Test User"
|
||||
assert "Administrators" in result.groups
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_token_invalid(self, auth_service):
|
||||
"""validate_token should raise ValueError for invalid token."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 401
|
||||
|
||||
with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid or expired token"):
|
||||
await auth_service.validate_token("invalid_token")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_token_service_unavailable(self, auth_service):
|
||||
"""validate_token should raise ValueError when service unavailable."""
|
||||
import httpx
|
||||
|
||||
with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
side_effect=httpx.RequestError("Connection failed")
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Authentication service unavailable"):
|
||||
await auth_service.validate_token("token")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# User Sync Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestSyncUser:
|
||||
"""Test user synchronization from OIDC token."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_user_creates_new_user(self, auth_service, sample_token_info, mock_session):
|
||||
"""sync_user should create new user when not found."""
|
||||
# Mock no existing user found
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = None
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
user, is_new = await auth_service.sync_user(sample_token_info)
|
||||
|
||||
assert is_new is True
|
||||
assert mock_session.add.call_count == 2 # User and Preferences
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_user_updates_existing_user(
|
||||
self, auth_service, sample_token_info, sample_user, mock_session
|
||||
):
|
||||
"""sync_user should update existing user when found."""
|
||||
# Mock existing user found
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = sample_user
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
# Update token info with matching authentik_id
|
||||
sample_token_info.sub = str(sample_user.authentik_id)
|
||||
|
||||
user, is_new = await auth_service.sync_user(sample_token_info)
|
||||
|
||||
assert is_new is False
|
||||
assert user.email == sample_token_info.email
|
||||
assert user.name == sample_token_info.name
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_user_updates_last_login(
|
||||
self, auth_service, sample_token_info, sample_user, mock_session
|
||||
):
|
||||
"""sync_user should update last_login timestamp."""
|
||||
old_login = sample_user.last_login
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = sample_user
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
sample_token_info.sub = str(sample_user.authentik_id)
|
||||
|
||||
user, _ = await auth_service.sync_user(sample_token_info)
|
||||
|
||||
assert user.last_login is not None
|
||||
# last_login should be updated (or same if happened in same second)
|
||||
assert user.last_login >= old_login or user.last_login is not None
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Role Sync Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestSyncRoles:
|
||||
"""Test role synchronization from Authentik groups via group_roles."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_roles_from_groups(
|
||||
self, auth_service, sample_user, sample_group, mock_session
|
||||
):
|
||||
"""sync_roles should get roles from matching groups."""
|
||||
# Mock finding groups with roles
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalars.return_value.all.return_value = [sample_group]
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
roles = await auth_service.sync_roles(sample_user, ["Administrators"])
|
||||
|
||||
assert len(roles) == 1
|
||||
assert roles[0].name == "control-room.general:admin"
|
||||
assert sample_user.roles == roles
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_roles_no_matching_groups(
|
||||
self, auth_service, sample_user, mock_session
|
||||
):
|
||||
"""sync_roles should return empty list when no groups match."""
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalars.return_value.all.return_value = []
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
roles = await auth_service.sync_roles(sample_user, ["NonExistentGroup"])
|
||||
|
||||
assert len(roles) == 0
|
||||
assert sample_user.roles == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_roles_deduplicates_roles(
|
||||
self, auth_service, sample_user, sample_role, mock_session
|
||||
):
|
||||
"""sync_roles should deduplicate roles from multiple groups."""
|
||||
# Create two groups with the same role
|
||||
group1 = Group(
|
||||
id=uuid.uuid4(),
|
||||
authentik_id=uuid.uuid4(),
|
||||
name="Group1",
|
||||
is_superuser=False,
|
||||
member_count=1,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
synced_at=datetime.now(timezone.utc),
|
||||
)
|
||||
group1.roles = [sample_role]
|
||||
|
||||
group2 = Group(
|
||||
id=uuid.uuid4(),
|
||||
authentik_id=uuid.uuid4(),
|
||||
name="Group2",
|
||||
is_superuser=False,
|
||||
member_count=1,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
synced_at=datetime.now(timezone.utc),
|
||||
)
|
||||
group2.roles = [sample_role] # Same role
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalars.return_value.all.return_value = [group1, group2]
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
roles = await auth_service.sync_roles(sample_user, ["Group1", "Group2"])
|
||||
|
||||
# Should only have one role despite appearing in two groups
|
||||
assert len(roles) == 1
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Schema Conversion Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestSchemaConversions:
|
||||
"""Test model to schema conversions."""
|
||||
|
||||
def test_user_to_schema(self, auth_service, sample_user):
|
||||
"""user_to_schema should convert User model to UserSchema."""
|
||||
schema = auth_service.user_to_schema(sample_user)
|
||||
|
||||
assert schema.id == sample_user.id
|
||||
assert schema.authentik_id == sample_user.authentik_id
|
||||
assert schema.email == sample_user.email
|
||||
assert schema.name == sample_user.name
|
||||
assert schema.avatar_url == sample_user.avatar_url
|
||||
|
||||
def test_roles_to_schema(self, auth_service, sample_role):
|
||||
"""roles_to_schema should convert Role models to RoleSchemas."""
|
||||
schemas = auth_service.roles_to_schema([sample_role])
|
||||
|
||||
assert len(schemas) == 1
|
||||
assert schemas[0].id == sample_role.id
|
||||
assert schemas[0].name == sample_role.name
|
||||
assert schemas[0].domain == sample_role.domain
|
||||
assert schemas[0].category == sample_role.category
|
||||
assert schemas[0].action == sample_role.action
|
||||
|
||||
def test_roles_to_schema_empty_list(self, auth_service):
|
||||
"""roles_to_schema should handle empty list."""
|
||||
schemas = auth_service.roles_to_schema([])
|
||||
assert schemas == []
|
||||
|
||||
def test_preferences_to_schema(self, auth_service, sample_preferences):
|
||||
"""preferences_to_schema should convert UserPreferences to schema."""
|
||||
schema = auth_service.preferences_to_schema(sample_preferences)
|
||||
|
||||
assert schema.theme == sample_preferences.theme
|
||||
assert schema.default_room == sample_preferences.default_room
|
||||
assert schema.preferences_json == sample_preferences.preferences_json
|
||||
|
||||
def test_preferences_to_schema_none(self, auth_service):
|
||||
"""preferences_to_schema should return defaults for None."""
|
||||
schema = auth_service.preferences_to_schema(None)
|
||||
|
||||
assert schema.theme == "system"
|
||||
assert schema.default_room == "front-hall"
|
||||
assert schema.preferences_json == {}
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# List Operations Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestListOperations:
|
||||
"""Test list operations for users, groups, and roles."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_users(self, auth_service, sample_user, mock_session):
|
||||
"""list_users should return paginated user list."""
|
||||
sample_user.roles = []
|
||||
|
||||
# Mock count query
|
||||
count_result = MagicMock()
|
||||
count_result.scalar.return_value = 1
|
||||
|
||||
# Mock users query
|
||||
users_result = MagicMock()
|
||||
users_result.scalars.return_value.all.return_value = [sample_user]
|
||||
|
||||
mock_session.execute.side_effect = [count_result, users_result]
|
||||
|
||||
items, total = await auth_service.list_users()
|
||||
|
||||
assert total == 1
|
||||
assert len(items) == 1
|
||||
assert items[0].email == sample_user.email
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_users_with_search(self, auth_service, mock_session):
|
||||
"""list_users should filter by search query."""
|
||||
count_result = MagicMock()
|
||||
count_result.scalar.return_value = 0
|
||||
|
||||
users_result = MagicMock()
|
||||
users_result.scalars.return_value.all.return_value = []
|
||||
|
||||
mock_session.execute.side_effect = [count_result, users_result]
|
||||
|
||||
items, total = await auth_service.list_users(search="nonexistent")
|
||||
|
||||
assert total == 0
|
||||
assert len(items) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_groups(self, auth_service, sample_group, mock_session):
|
||||
"""list_groups should return paginated group list with roles."""
|
||||
count_result = MagicMock()
|
||||
count_result.scalar.return_value = 1
|
||||
|
||||
groups_result = MagicMock()
|
||||
groups_result.scalars.return_value.all.return_value = [sample_group]
|
||||
|
||||
mock_session.execute.side_effect = [count_result, groups_result]
|
||||
|
||||
items, total = await auth_service.list_groups()
|
||||
|
||||
assert total == 1
|
||||
assert len(items) == 1
|
||||
assert items[0].name == sample_group.name
|
||||
assert len(items[0].roles) == 1 # Should include role names
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_roles(self, auth_service, sample_role, mock_session):
|
||||
"""list_roles should return all roles ordered by domain."""
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalars.return_value.all.return_value = [sample_role]
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
roles = await auth_service.list_roles()
|
||||
|
||||
assert len(roles) == 1
|
||||
assert roles[0].name == sample_role.name
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Group-Role Management Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestGroupRoleManagement:
|
||||
"""Test group-role assignment and removal."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_group_by_id(self, auth_service, sample_group, mock_session):
|
||||
"""get_group_by_id should return group with roles."""
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = sample_group
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
group = await auth_service.get_group_by_id(sample_group.id)
|
||||
|
||||
assert group is not None
|
||||
assert group.id == sample_group.id
|
||||
assert len(group.roles) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_group_by_id_not_found(self, auth_service, mock_session):
|
||||
"""get_group_by_id should return None when not found."""
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = None
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
group = await auth_service.get_group_by_id(uuid.uuid4())
|
||||
|
||||
assert group is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_assign_role_to_group(
|
||||
self, auth_service, sample_group, sample_role, mock_session
|
||||
):
|
||||
"""assign_role_to_group should add role to group."""
|
||||
# Clear existing roles for this test
|
||||
sample_group.roles = []
|
||||
|
||||
# Mock group lookup
|
||||
group_result = MagicMock()
|
||||
group_result.scalar_one_or_none.return_value = sample_group
|
||||
|
||||
# Mock role lookup
|
||||
role_result = MagicMock()
|
||||
role_result.scalar_one_or_none.return_value = sample_role
|
||||
|
||||
mock_session.execute.side_effect = [group_result, role_result]
|
||||
|
||||
group = await auth_service.assign_role_to_group(sample_group.id, sample_role.id)
|
||||
|
||||
assert sample_role in group.roles
|
||||
mock_session.flush.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_assign_role_to_group_already_assigned(
|
||||
self, auth_service, sample_group, sample_role, mock_session
|
||||
):
|
||||
"""assign_role_to_group should not duplicate if already assigned."""
|
||||
# Group already has this role
|
||||
sample_group.roles = [sample_role]
|
||||
original_count = len(sample_group.roles)
|
||||
|
||||
group_result = MagicMock()
|
||||
group_result.scalar_one_or_none.return_value = sample_group
|
||||
|
||||
role_result = MagicMock()
|
||||
role_result.scalar_one_or_none.return_value = sample_role
|
||||
|
||||
mock_session.execute.side_effect = [group_result, role_result]
|
||||
|
||||
group = await auth_service.assign_role_to_group(sample_group.id, sample_role.id)
|
||||
|
||||
assert len(group.roles) == original_count # No duplicate
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_assign_role_to_group_group_not_found(self, auth_service, mock_session):
|
||||
"""assign_role_to_group should raise ValueError when group not found."""
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = None
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
with pytest.raises(ValueError, match="Group not found"):
|
||||
await auth_service.assign_role_to_group(uuid.uuid4(), uuid.uuid4())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_assign_role_to_group_role_not_found(
|
||||
self, auth_service, sample_group, mock_session
|
||||
):
|
||||
"""assign_role_to_group should raise ValueError when role not found."""
|
||||
sample_group.roles = []
|
||||
|
||||
group_result = MagicMock()
|
||||
group_result.scalar_one_or_none.return_value = sample_group
|
||||
|
||||
role_result = MagicMock()
|
||||
role_result.scalar_one_or_none.return_value = None
|
||||
|
||||
mock_session.execute.side_effect = [group_result, role_result]
|
||||
|
||||
with pytest.raises(ValueError, match="Role not found"):
|
||||
await auth_service.assign_role_to_group(sample_group.id, uuid.uuid4())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_role_from_group(
|
||||
self, auth_service, sample_group, sample_role, mock_session
|
||||
):
|
||||
"""remove_role_from_group should remove role from group."""
|
||||
# Group has this role
|
||||
sample_group.roles = [sample_role]
|
||||
|
||||
group_result = MagicMock()
|
||||
group_result.scalar_one_or_none.return_value = sample_group
|
||||
|
||||
role_result = MagicMock()
|
||||
role_result.scalar_one_or_none.return_value = sample_role
|
||||
|
||||
mock_session.execute.side_effect = [group_result, role_result]
|
||||
|
||||
group = await auth_service.remove_role_from_group(sample_group.id, sample_role.id)
|
||||
|
||||
assert sample_role not in group.roles
|
||||
mock_session.flush.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_role_from_group_not_assigned(
|
||||
self, auth_service, sample_group, sample_role, mock_session
|
||||
):
|
||||
"""remove_role_from_group should handle role not assigned gracefully."""
|
||||
# Group does not have this role
|
||||
sample_group.roles = []
|
||||
|
||||
group_result = MagicMock()
|
||||
group_result.scalar_one_or_none.return_value = sample_group
|
||||
|
||||
role_result = MagicMock()
|
||||
role_result.scalar_one_or_none.return_value = sample_role
|
||||
|
||||
mock_session.execute.side_effect = [group_result, role_result]
|
||||
|
||||
group = await auth_service.remove_role_from_group(sample_group.id, sample_role.id)
|
||||
|
||||
# Should complete without error
|
||||
assert len(group.roles) == 0
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Role Schema Tests
|
||||
# =============================================================================
|
||||
|
||||
class TestRoleSchema:
|
||||
"""Test RoleSchema validation."""
|
||||
|
||||
def test_role_schema_creation(self):
|
||||
"""RoleSchema should be creatable with valid data."""
|
||||
schema = RoleSchema(
|
||||
id=uuid.uuid4(),
|
||||
name="control-room.general:admin",
|
||||
domain="control-room",
|
||||
category="general",
|
||||
action="admin",
|
||||
)
|
||||
|
||||
assert schema.name == "control-room.general:admin"
|
||||
assert schema.domain == "control-room"
|
||||
assert schema.category == "general"
|
||||
assert schema.action == "admin"
|
||||
|
||||
def test_role_schema_category_default(self):
|
||||
"""RoleSchema should default category to 'general'."""
|
||||
schema = RoleSchema(
|
||||
id=uuid.uuid4(),
|
||||
name="media.general:viewer",
|
||||
domain="media",
|
||||
action="viewer",
|
||||
)
|
||||
|
||||
assert schema.category == "general"
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Tests for config module."""
|
||||
import pytest
|
||||
from src.config import (
|
||||
__version__,
|
||||
Settings,
|
||||
get_settings,
|
||||
_get_version_from_pyproject,
|
||||
)
|
||||
|
||||
|
||||
class TestVersion:
|
||||
"""Test version loading from pyproject.toml."""
|
||||
|
||||
def test_version_is_loaded(self):
|
||||
"""Version should be loaded from pyproject.toml."""
|
||||
assert __version__ is not None
|
||||
assert isinstance(__version__, str)
|
||||
|
||||
def test_version_format(self):
|
||||
"""Version should follow semver format."""
|
||||
parts = __version__.split(".")
|
||||
assert len(parts) >= 2, "Version should have at least major.minor"
|
||||
assert all(p.isdigit() for p in parts), "Version parts should be numeric"
|
||||
|
||||
def test_version_matches_settings(self):
|
||||
"""Settings app_version should match module version."""
|
||||
settings = get_settings()
|
||||
assert settings.app_version == __version__
|
||||
|
||||
|
||||
class TestGetVersionFromPyproject:
|
||||
"""Test the version loading function."""
|
||||
|
||||
def test_returns_string(self):
|
||||
"""Should return a string version."""
|
||||
version = _get_version_from_pyproject()
|
||||
assert isinstance(version, str)
|
||||
|
||||
def test_returns_valid_version(self):
|
||||
"""Should return a valid version (not 0.0.0 if file exists)."""
|
||||
version = _get_version_from_pyproject()
|
||||
# Since pyproject.toml exists, version should not be fallback
|
||||
assert version != "0.0.0"
|
||||
|
||||
|
||||
class TestSettings:
|
||||
"""Test Settings configuration class."""
|
||||
|
||||
def test_settings_has_app_name(self):
|
||||
"""Settings should have app_name."""
|
||||
settings = get_settings()
|
||||
assert settings.app_name == "Core Code API"
|
||||
|
||||
def test_settings_has_version(self):
|
||||
"""Settings should have app_version."""
|
||||
settings = get_settings()
|
||||
assert settings.app_version is not None
|
||||
|
||||
def test_settings_default_host(self):
|
||||
"""Settings should have default host."""
|
||||
settings = get_settings()
|
||||
assert settings.host == "0.0.0.0"
|
||||
|
||||
def test_settings_default_port(self):
|
||||
"""Settings should have default port."""
|
||||
settings = get_settings()
|
||||
assert settings.port == 8083
|
||||
|
||||
def test_no_kuma_settings(self):
|
||||
"""Settings should not have Kuma-related attributes."""
|
||||
settings = get_settings()
|
||||
assert not hasattr(settings, "kuma_url")
|
||||
assert not hasattr(settings, "kuma_username")
|
||||
assert not hasattr(settings, "kuma_password")
|
||||
assert not hasattr(settings, "kuma_api_key")
|
||||
|
||||
|
||||
class TestGetSettings:
|
||||
"""Test get_settings function."""
|
||||
|
||||
def test_returns_settings_instance(self):
|
||||
"""Should return a Settings instance."""
|
||||
settings = get_settings()
|
||||
assert isinstance(settings, Settings)
|
||||
|
||||
def test_returns_cached_instance(self):
|
||||
"""Should return the same cached instance."""
|
||||
settings1 = get_settings()
|
||||
settings2 = get_settings()
|
||||
assert settings1 is settings2
|
||||
|
||||
def test_model_aliases_property(self):
|
||||
"""Model aliases property should return dict."""
|
||||
settings = get_settings()
|
||||
aliases = settings.model_aliases
|
||||
assert isinstance(aliases, dict)
|
||||
assert "gpt-3.5-turbo" in aliases
|
||||
assert "gpt-4" in aliases
|
||||
@@ -0,0 +1,199 @@
|
||||
"""Tests for dashboard endpoints registration and OpenAPI spec."""
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.main import app
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
"""Create a test client."""
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
class TestDashboardOpenAPISpec:
|
||||
"""Test that dashboard endpoints are documented in OpenAPI spec."""
|
||||
|
||||
def test_quick_links_list_in_openapi(self, client):
|
||||
"""Quick links list endpoint should be in OpenAPI spec."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
spec = response.json()
|
||||
assert "/dashboard/quick-links" in spec["paths"]
|
||||
|
||||
def test_quick_links_get_in_openapi(self, client):
|
||||
"""Quick links get endpoint should be in OpenAPI spec."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
spec = response.json()
|
||||
assert "/dashboard/quick-links/{link_id}" in spec["paths"]
|
||||
|
||||
def test_quick_links_reorder_in_openapi(self, client):
|
||||
"""Quick links reorder endpoint should be in OpenAPI spec."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
spec = response.json()
|
||||
assert "/dashboard/quick-links/reorder" in spec["paths"]
|
||||
|
||||
def test_widgets_list_in_openapi(self, client):
|
||||
"""Widgets list endpoint should be in OpenAPI spec."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
spec = response.json()
|
||||
assert "/dashboard/widgets" in spec["paths"]
|
||||
|
||||
def test_widgets_get_in_openapi(self, client):
|
||||
"""Widgets get endpoint should be in OpenAPI spec."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
spec = response.json()
|
||||
assert "/dashboard/widgets/{widget_id}" in spec["paths"]
|
||||
|
||||
def test_quick_links_supports_crud_operations(self, client):
|
||||
"""Quick links should support all CRUD operations."""
|
||||
response = client.get("/openapi.json")
|
||||
spec = response.json()
|
||||
|
||||
# List endpoint
|
||||
list_path = spec["paths"].get("/dashboard/quick-links", {})
|
||||
assert "get" in list_path # List
|
||||
assert "post" in list_path # Create
|
||||
|
||||
# Item endpoint
|
||||
item_path = spec["paths"].get("/dashboard/quick-links/{link_id}", {})
|
||||
assert "get" in item_path # Read
|
||||
assert "put" in item_path # Update
|
||||
assert "delete" in item_path # Delete
|
||||
|
||||
def test_widgets_supports_crud_operations(self, client):
|
||||
"""Widgets should support all CRUD operations."""
|
||||
response = client.get("/openapi.json")
|
||||
spec = response.json()
|
||||
|
||||
# List endpoint
|
||||
list_path = spec["paths"].get("/dashboard/widgets", {})
|
||||
assert "get" in list_path # List
|
||||
assert "post" in list_path # Create
|
||||
|
||||
# Item endpoint
|
||||
item_path = spec["paths"].get("/dashboard/widgets/{widget_id}", {})
|
||||
assert "get" in item_path # Read
|
||||
assert "put" in item_path # Update
|
||||
assert "delete" in item_path # Delete
|
||||
|
||||
|
||||
class TestDashboardSchemaValidation:
|
||||
"""Test that request validation works correctly."""
|
||||
|
||||
def test_create_quick_link_requires_title(self, client):
|
||||
"""Create quick link should require title (422 for validation)."""
|
||||
response = client.post(
|
||||
"/dashboard/quick-links",
|
||||
json={
|
||||
"url": "https://example.com",
|
||||
},
|
||||
)
|
||||
# Either 422 for validation or 401/403/500 for auth
|
||||
assert response.status_code in [401, 403, 422, 500]
|
||||
|
||||
def test_create_widget_requires_widget_type(self, client):
|
||||
"""Create widget should require widget_type (422 for validation)."""
|
||||
response = client.post(
|
||||
"/dashboard/widgets",
|
||||
json={},
|
||||
)
|
||||
assert response.status_code in [401, 403, 422, 500]
|
||||
|
||||
def test_reorder_requires_link_ids(self, client):
|
||||
"""Reorder should require link_ids list (422 for validation)."""
|
||||
response = client.post(
|
||||
"/dashboard/quick-links/reorder",
|
||||
json={},
|
||||
)
|
||||
assert response.status_code in [401, 403, 422, 500]
|
||||
|
||||
|
||||
class TestDashboardControllerInit:
|
||||
"""Test dashboard controller initialization."""
|
||||
|
||||
def test_controller_module_imports(self):
|
||||
"""Dashboard controller should be importable."""
|
||||
from src.domains.dashboard.controller import DashboardController, dashboard_controller
|
||||
assert DashboardController is not None
|
||||
assert dashboard_controller is not None
|
||||
|
||||
def test_controller_has_correct_prefix(self):
|
||||
"""Dashboard controller should have correct prefix."""
|
||||
from src.domains.dashboard.controller import dashboard_controller
|
||||
assert dashboard_controller.prefix == "/dashboard"
|
||||
|
||||
def test_controller_has_correct_tags(self):
|
||||
"""Dashboard controller should have correct tags."""
|
||||
from src.domains.dashboard.controller import dashboard_controller
|
||||
assert "Dashboard" in dashboard_controller.tags
|
||||
|
||||
|
||||
class TestDashboardServiceInit:
|
||||
"""Test dashboard service initialization."""
|
||||
|
||||
def test_service_module_imports(self):
|
||||
"""Dashboard service should be importable."""
|
||||
from src.domains.dashboard.service import DashboardService, get_dashboard_service
|
||||
assert DashboardService is not None
|
||||
assert get_dashboard_service is not None
|
||||
|
||||
def test_service_singleton(self):
|
||||
"""get_dashboard_service should return singleton."""
|
||||
from src.domains.dashboard.service import get_dashboard_service
|
||||
|
||||
service1 = get_dashboard_service()
|
||||
service2 = get_dashboard_service()
|
||||
assert service1 is service2
|
||||
|
||||
|
||||
class TestDashboardModels:
|
||||
"""Test dashboard models."""
|
||||
|
||||
def test_quick_link_model_imports(self):
|
||||
"""QuickLink model should be importable."""
|
||||
from src.domains.dashboard.models import QuickLink
|
||||
assert QuickLink is not None
|
||||
|
||||
def test_dashboard_widget_model_imports(self):
|
||||
"""DashboardWidget model should be importable."""
|
||||
from src.domains.dashboard.models import DashboardWidget
|
||||
assert DashboardWidget is not None
|
||||
|
||||
|
||||
class TestDashboardSchemas:
|
||||
"""Test dashboard schemas."""
|
||||
|
||||
def test_quick_link_schemas_import(self):
|
||||
"""QuickLink schemas should be importable."""
|
||||
from src.domains.dashboard.schemas import (
|
||||
QuickLinkCreate,
|
||||
QuickLinkUpdate,
|
||||
QuickLinkResponse,
|
||||
QuickLinkListResponse,
|
||||
QuickLinkReorderRequest,
|
||||
QuickLinkReorderResponse,
|
||||
)
|
||||
assert QuickLinkCreate is not None
|
||||
assert QuickLinkUpdate is not None
|
||||
assert QuickLinkResponse is not None
|
||||
assert QuickLinkListResponse is not None
|
||||
assert QuickLinkReorderRequest is not None
|
||||
assert QuickLinkReorderResponse is not None
|
||||
|
||||
def test_dashboard_widget_schemas_import(self):
|
||||
"""DashboardWidget schemas should be importable."""
|
||||
from src.domains.dashboard.schemas import (
|
||||
DashboardWidgetCreate,
|
||||
DashboardWidgetUpdate,
|
||||
DashboardWidgetResponse,
|
||||
DashboardWidgetListResponse,
|
||||
)
|
||||
assert DashboardWidgetCreate is not None
|
||||
assert DashboardWidgetUpdate is not None
|
||||
assert DashboardWidgetResponse is not None
|
||||
assert DashboardWidgetListResponse is not None
|
||||
@@ -0,0 +1,331 @@
|
||||
"""Tests for DNS service."""
|
||||
import pytest
|
||||
from unittest.mock import patch, MagicMock
|
||||
import dns.resolver
|
||||
import dns.exception
|
||||
|
||||
from src.domains.tools.dns.service import DNSService
|
||||
from src.domains.tools.dns.schemas import DNSLookupRequest, DNSRecord
|
||||
from src.domains.tools.dns.exceptions import DNSQueryError
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def dns_service():
|
||||
"""Create a DNSService instance."""
|
||||
return DNSService()
|
||||
|
||||
|
||||
class TestDNSServiceInit:
|
||||
"""Test DNSService initialization."""
|
||||
|
||||
def test_service_has_resolver(self, dns_service):
|
||||
"""Service should have resolver configured."""
|
||||
assert dns_service.resolver is not None
|
||||
|
||||
def test_service_has_timeout(self, dns_service):
|
||||
"""Service should have timeout configured."""
|
||||
assert dns_service.resolver.timeout == 5.0
|
||||
assert dns_service.resolver.lifetime == 10.0
|
||||
|
||||
def test_supported_record_types(self, dns_service):
|
||||
"""Service should have supported record types."""
|
||||
assert "A" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
assert "AAAA" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
assert "MX" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
assert "TXT" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
assert "CNAME" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
assert "NS" in dns_service.SUPPORTED_RECORD_TYPES
|
||||
|
||||
|
||||
class TestDNSServiceLookup:
|
||||
"""Test DNS lookup functionality."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_validates_record_type(self, dns_service):
|
||||
"""lookup should raise error for unsupported record type."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="INVALID")
|
||||
|
||||
with pytest.raises(DNSQueryError) as exc_info:
|
||||
await dns_service.lookup(request)
|
||||
|
||||
assert "Unsupported record type" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_returns_response_on_success(self, dns_service):
|
||||
"""lookup should return DNSLookupResponse on success."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="A")
|
||||
|
||||
# Mock the resolver
|
||||
mock_answer = MagicMock()
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.__str__ = MagicMock(return_value="93.184.216.34")
|
||||
mock_answer.__iter__ = MagicMock(return_value=iter([mock_rdata]))
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.return_value = mock_answer
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is True
|
||||
assert response.domain == "example.com"
|
||||
assert response.record_type == "A"
|
||||
assert len(response.records) > 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_uses_custom_nameserver(self, dns_service):
|
||||
"""lookup should use custom nameserver when specified."""
|
||||
request = DNSLookupRequest(
|
||||
domain="example.com",
|
||||
record_type="A",
|
||||
nameserver="1.1.1.1"
|
||||
)
|
||||
|
||||
mock_answer = MagicMock()
|
||||
mock_answer.__iter__ = MagicMock(return_value=iter([]))
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.return_value = mock_answer
|
||||
mock_resolver.nameservers = []
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
# Verify nameserver was set
|
||||
assert mock_resolver.nameservers == ["1.1.1.1"]
|
||||
assert response.nameserver_used == "1.1.1.1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_handles_nxdomain(self, dns_service):
|
||||
"""lookup should handle NXDOMAIN (domain not found)."""
|
||||
request = DNSLookupRequest(domain="nonexistent.invalid", record_type="A")
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.side_effect = dns.resolver.NXDOMAIN()
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is False
|
||||
assert "Domain not found" in response.error_message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_handles_no_answer(self, dns_service):
|
||||
"""lookup should handle NoAnswer (no records of type)."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="AAAA")
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.side_effect = dns.resolver.NoAnswer()
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is False
|
||||
assert "No AAAA records found" in response.error_message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_handles_timeout(self, dns_service):
|
||||
"""lookup should handle DNS timeout."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="A")
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.side_effect = dns.resolver.Timeout()
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is False
|
||||
assert "timeout" in response.error_message.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_handles_dns_exception(self, dns_service):
|
||||
"""lookup should handle generic DNS exceptions."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="A")
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.side_effect = dns.exception.DNSException("DNS error")
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is False
|
||||
assert "DNS error" in response.error_message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_handles_unexpected_exception(self, dns_service):
|
||||
"""lookup should handle unexpected exceptions."""
|
||||
request = DNSLookupRequest(domain="example.com", record_type="A")
|
||||
|
||||
with patch("dns.resolver.Resolver") as mock_resolver_class:
|
||||
mock_resolver = MagicMock()
|
||||
mock_resolver.resolve.side_effect = Exception("Unexpected error")
|
||||
mock_resolver.nameservers = ["8.8.8.8"]
|
||||
mock_resolver_class.return_value = mock_resolver
|
||||
|
||||
response = await dns_service.lookup(request)
|
||||
|
||||
assert response.success is False
|
||||
assert "Unexpected error" in response.error_message
|
||||
|
||||
|
||||
class TestDNSServiceParseRecord:
|
||||
"""Test record parsing."""
|
||||
|
||||
def test_parse_a_record(self, dns_service):
|
||||
"""_parse_record should parse A record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.__str__ = MagicMock(return_value="192.168.1.1")
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "A")
|
||||
|
||||
assert record is not None
|
||||
assert record.value == "192.168.1.1"
|
||||
|
||||
def test_parse_aaaa_record(self, dns_service):
|
||||
"""_parse_record should parse AAAA record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.__str__ = MagicMock(return_value="2001:db8::1")
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "AAAA")
|
||||
|
||||
assert record is not None
|
||||
assert record.value == "2001:db8::1"
|
||||
|
||||
def test_parse_mx_record(self, dns_service):
|
||||
"""_parse_record should parse MX record with priority."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.exchange = "mail.example.com"
|
||||
mock_rdata.preference = 10
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "MX")
|
||||
|
||||
assert record is not None
|
||||
assert "mail.example.com" in record.value
|
||||
assert record.priority == 10
|
||||
|
||||
def test_parse_txt_record(self, dns_service):
|
||||
"""_parse_record should parse TXT record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.strings = [b"v=spf1 include:_spf.google.com ~all"]
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "TXT")
|
||||
|
||||
assert record is not None
|
||||
assert "spf1" in record.value
|
||||
|
||||
def test_parse_cname_record(self, dns_service):
|
||||
"""_parse_record should parse CNAME record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.target = "alias.example.com"
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "CNAME")
|
||||
|
||||
assert record is not None
|
||||
assert "alias.example.com" in record.value
|
||||
|
||||
def test_parse_ns_record(self, dns_service):
|
||||
"""_parse_record should parse NS record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.target = "ns1.example.com"
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "NS")
|
||||
|
||||
assert record is not None
|
||||
assert "ns1.example.com" in record.value
|
||||
|
||||
def test_parse_soa_record(self, dns_service):
|
||||
"""_parse_record should parse SOA record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.mname = "ns1.example.com"
|
||||
mock_rdata.rname = "admin.example.com"
|
||||
mock_rdata.serial = 2024010101
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "SOA")
|
||||
|
||||
assert record is not None
|
||||
assert "ns1.example.com" in record.value
|
||||
assert "2024010101" in record.value
|
||||
|
||||
def test_parse_srv_record(self, dns_service):
|
||||
"""_parse_record should parse SRV record."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.target = "server.example.com"
|
||||
mock_rdata.port = 443
|
||||
mock_rdata.priority = 10
|
||||
mock_rdata.weight = 100
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "SRV")
|
||||
|
||||
assert record is not None
|
||||
assert "server.example.com" in record.value
|
||||
assert "port=443" in record.value
|
||||
assert record.priority == 10
|
||||
|
||||
def test_parse_record_returns_none_on_error(self, dns_service):
|
||||
"""_parse_record should return None on parsing error."""
|
||||
mock_rdata = MagicMock()
|
||||
mock_rdata.__str__ = MagicMock(side_effect=Exception("Parse error"))
|
||||
|
||||
record = dns_service._parse_record(mock_rdata, "A")
|
||||
|
||||
assert record is None
|
||||
|
||||
|
||||
class TestDNSServiceErrorResponse:
|
||||
"""Test error response generation."""
|
||||
|
||||
def test_error_response_includes_domain(self, dns_service):
|
||||
"""_error_response should include domain."""
|
||||
import time
|
||||
request = DNSLookupRequest(domain="test.example.com", record_type="A")
|
||||
|
||||
response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error")
|
||||
|
||||
assert response.domain == "test.example.com"
|
||||
|
||||
def test_error_response_includes_record_type(self, dns_service):
|
||||
"""_error_response should include record type."""
|
||||
import time
|
||||
request = DNSLookupRequest(domain="test.example.com", record_type="mx")
|
||||
|
||||
response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error")
|
||||
|
||||
assert response.record_type == "MX" # Should be uppercase
|
||||
|
||||
def test_error_response_has_empty_records(self, dns_service):
|
||||
"""_error_response should have empty records list."""
|
||||
import time
|
||||
request = DNSLookupRequest(domain="test.example.com", record_type="A")
|
||||
|
||||
response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error")
|
||||
|
||||
assert response.records == []
|
||||
|
||||
def test_error_response_has_success_false(self, dns_service):
|
||||
"""_error_response should have success=False."""
|
||||
import time
|
||||
request = DNSLookupRequest(domain="test.example.com", record_type="A")
|
||||
|
||||
response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error")
|
||||
|
||||
assert response.success is False
|
||||
|
||||
def test_error_response_includes_error_message(self, dns_service):
|
||||
"""_error_response should include error message."""
|
||||
import time
|
||||
request = DNSLookupRequest(domain="test.example.com", record_type="A")
|
||||
|
||||
response = dns_service._error_response(request, "8.8.8.8", time.time(), "Specific error")
|
||||
|
||||
assert response.error_message == "Specific error"
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Tests for health endpoints."""
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from unittest.mock import patch, AsyncMock
|
||||
|
||||
from src.main import app
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
"""Create a test client."""
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
class TestRootEndpoint:
|
||||
"""Test root endpoint."""
|
||||
|
||||
def test_root_returns_200(self, client):
|
||||
"""Root endpoint should return 200."""
|
||||
response = client.get("/")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_root_returns_service_info(self, client):
|
||||
"""Root endpoint should return service information."""
|
||||
response = client.get("/")
|
||||
data = response.json()
|
||||
|
||||
assert "service" in data
|
||||
assert "version" in data
|
||||
assert "status" in data
|
||||
assert data["service"] == "Core Code API"
|
||||
assert data["status"] == "healthy"
|
||||
|
||||
def test_root_returns_docs_link(self, client):
|
||||
"""Root endpoint should return docs link."""
|
||||
response = client.get("/")
|
||||
data = response.json()
|
||||
|
||||
assert "docs" in data
|
||||
assert data["docs"] == "/docs"
|
||||
|
||||
|
||||
class TestHealthEndpoint:
|
||||
"""Test /health endpoint."""
|
||||
|
||||
def test_health_returns_200(self, client):
|
||||
"""Health endpoint should return 200."""
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_health_returns_status(self, client):
|
||||
"""Health endpoint should return status information."""
|
||||
response = client.get("/health")
|
||||
data = response.json()
|
||||
|
||||
assert "status" in data
|
||||
assert "version" in data
|
||||
assert data["status"] == "healthy"
|
||||
|
||||
|
||||
class TestOpenAPIEndpoint:
|
||||
"""Test OpenAPI documentation endpoints."""
|
||||
|
||||
def test_openapi_spec_available(self, client):
|
||||
"""OpenAPI spec should be available."""
|
||||
response = client.get("/openapi.json")
|
||||
assert response.status_code == 200
|
||||
|
||||
data = response.json()
|
||||
assert "openapi" in data
|
||||
assert "info" in data
|
||||
|
||||
def test_swagger_ui_available(self, client):
|
||||
"""Swagger UI should be available."""
|
||||
response = client.get("/docs")
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
class TestFullHealthCheck:
|
||||
"""Test /health/full endpoint."""
|
||||
|
||||
@patch("src.shared.database.Database.health_check")
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_full_health_returns_503_when_unhealthy(self, mock_get_ollama, mock_db_health, client):
|
||||
"""Full health should return 503 when Ollama unhealthy."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = False
|
||||
mock_get_ollama.return_value = mock_client
|
||||
mock_db_health.return_value = True
|
||||
|
||||
response = client.get("/health/full")
|
||||
assert response.status_code == 503
|
||||
|
||||
@patch("src.shared.database.Database.health_check")
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_full_health_returns_components_status(self, mock_get_ollama, mock_db_health, client):
|
||||
"""Full health should return component status."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = False
|
||||
mock_get_ollama.return_value = mock_client
|
||||
mock_db_health.return_value = True
|
||||
|
||||
response = client.get("/health/full")
|
||||
data = response.json()
|
||||
|
||||
assert "status" in data
|
||||
assert "components" in data
|
||||
assert "ollama" in data["components"]
|
||||
assert "database" in data["components"]
|
||||
assert "response_time_ms" in data
|
||||
|
||||
@patch("src.shared.database.Database.health_check")
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_full_health_handles_list_models_error(self, mock_get_ollama, mock_db_health, client):
|
||||
"""Full health should handle list_models errors."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_client.list_models.side_effect = Exception("Connection error")
|
||||
mock_get_ollama.return_value = mock_client
|
||||
mock_db_health.return_value = True
|
||||
|
||||
response = client.get("/health/full")
|
||||
data = response.json()
|
||||
|
||||
# Should report error in component status
|
||||
assert "ollama" in data["components"]
|
||||
|
||||
@patch("src.shared.database.Database.health_check")
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_full_health_handles_health_check_exception(self, mock_get_ollama, mock_db_health, client):
|
||||
"""Full health should handle health check exceptions gracefully."""
|
||||
mock_client = AsyncMock()
|
||||
# Return False instead of raising exception to test unhealthy path
|
||||
mock_client.health_check.return_value = False
|
||||
mock_get_ollama.return_value = mock_client
|
||||
mock_db_health.return_value = True
|
||||
|
||||
response = client.get("/health/full")
|
||||
# Should return 503 for unhealthy
|
||||
assert response.status_code == 503
|
||||
data = response.json()
|
||||
assert data["status"] == "unhealthy"
|
||||
|
||||
|
||||
class TestDiagnosticsEndpoint:
|
||||
"""Test /health/diagnostics endpoint."""
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_returns_200(self, mock_get_ollama, client):
|
||||
"""Diagnostics should return 200."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
assert response.status_code == 200
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_returns_service_info(self, mock_get_ollama, client):
|
||||
"""Diagnostics should return service information."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
data = response.json()
|
||||
|
||||
assert "service" in data
|
||||
assert "name" in data["service"]
|
||||
assert "version" in data["service"]
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_returns_components(self, mock_get_ollama, client):
|
||||
"""Diagnostics should return component details."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
data = response.json()
|
||||
|
||||
assert "components" in data
|
||||
assert "ollama" in data["components"]
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_returns_configuration(self, mock_get_ollama, client):
|
||||
"""Diagnostics should return configuration info."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
data = response.json()
|
||||
|
||||
assert "configuration" in data
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_returns_response_time(self, mock_get_ollama, client):
|
||||
"""Diagnostics should return response time."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.return_value = True
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
data = response.json()
|
||||
|
||||
assert "response_time_ms" in data
|
||||
assert isinstance(data["response_time_ms"], int)
|
||||
|
||||
@patch("src.models.ollama_client.get_ollama_client")
|
||||
def test_diagnostics_handles_ollama_error(self, mock_get_ollama, client):
|
||||
"""Diagnostics should handle Ollama connection errors."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.health_check.side_effect = Exception("Connection refused")
|
||||
mock_get_ollama.return_value = mock_client
|
||||
|
||||
response = client.get("/health/diagnostics")
|
||||
data = response.json()
|
||||
|
||||
# Should still return 200 with error info
|
||||
assert response.status_code == 200
|
||||
assert "error" in data["components"]["ollama"]
|
||||
@@ -0,0 +1,471 @@
|
||||
"""Tests for Home Assistant client."""
|
||||
import pytest
|
||||
from unittest.mock import patch, AsyncMock, MagicMock
|
||||
import httpx
|
||||
|
||||
from src.shared.clients.homeassistant_client import HomeAssistantClient, get_homeassistant_client
|
||||
|
||||
|
||||
class TestHomeAssistantClientInit:
|
||||
"""Test HomeAssistantClient initialization."""
|
||||
|
||||
@patch("src.shared.clients.homeassistant_client.settings")
|
||||
def test_uses_settings_defaults(self, mock_settings):
|
||||
"""Client should use settings for defaults."""
|
||||
mock_settings.homeassistant_url = "http://ha.local:8123"
|
||||
mock_settings.homeassistant_token = "test_token"
|
||||
|
||||
client = HomeAssistantClient()
|
||||
|
||||
assert client.base_url == "http://ha.local:8123"
|
||||
assert client.token == "test_token"
|
||||
|
||||
def test_accepts_custom_url_and_token(self):
|
||||
"""Client should accept custom URL and token."""
|
||||
client = HomeAssistantClient(
|
||||
base_url="http://custom:8123",
|
||||
token="custom_token"
|
||||
)
|
||||
|
||||
assert client.base_url == "http://custom:8123"
|
||||
assert client.token == "custom_token"
|
||||
|
||||
def test_strips_trailing_slash_from_url(self):
|
||||
"""Client should strip trailing slash from URL."""
|
||||
client = HomeAssistantClient(
|
||||
base_url="http://custom:8123/",
|
||||
token="token"
|
||||
)
|
||||
|
||||
assert client.base_url == "http://custom:8123"
|
||||
|
||||
@patch("src.shared.clients.homeassistant_client.logger")
|
||||
@patch("src.shared.clients.homeassistant_client.settings")
|
||||
def test_warns_when_token_missing(self, mock_settings, mock_logger):
|
||||
"""Client should warn when token is not configured."""
|
||||
mock_settings.homeassistant_url = "http://ha.local:8123"
|
||||
mock_settings.homeassistant_token = ""
|
||||
|
||||
HomeAssistantClient()
|
||||
|
||||
mock_logger.warning.assert_called_once()
|
||||
|
||||
|
||||
class TestHomeAssistantClientHeaders:
|
||||
"""Test header generation."""
|
||||
|
||||
def test_get_headers_includes_bearer_token(self):
|
||||
"""Headers should include Bearer token."""
|
||||
client = HomeAssistantClient(
|
||||
base_url="http://ha:8123",
|
||||
token="my_token"
|
||||
)
|
||||
|
||||
headers = client._get_headers()
|
||||
|
||||
assert headers["Authorization"] == "Bearer my_token"
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
|
||||
class TestHomeAssistantClientHealthCheck:
|
||||
"""Test health check functionality."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_returns_healthy(self):
|
||||
"""Health check should return healthy when HA responds."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"version": "2024.12.0"}
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.health_check()
|
||||
|
||||
assert result["status"] == "healthy"
|
||||
assert result["connected"] is True
|
||||
assert result["platform"] == "home_assistant"
|
||||
assert result["version"] == "2024.12.0"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_returns_unhealthy_on_error(self):
|
||||
"""Health check should return unhealthy on connection error."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.side_effect = Exception("Connection refused")
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert result["connected"] is False
|
||||
assert "error" in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_returns_unhealthy_on_non_200(self):
|
||||
"""Health check should return unhealthy on non-200 status."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 401
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert result["connected"] is False
|
||||
|
||||
|
||||
class TestHomeAssistantClientStates:
|
||||
"""Test state retrieval methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_states_returns_list(self):
|
||||
"""get_states should return list of states."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
states = [
|
||||
{"entity_id": "light.test", "state": "on"},
|
||||
{"entity_id": "switch.test", "state": "off"}
|
||||
]
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = states
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.get_states()
|
||||
|
||||
assert result == states
|
||||
assert len(result) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_state_returns_single_entity(self):
|
||||
"""get_state should return single entity state."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
state = {"entity_id": "light.test", "state": "on", "attributes": {}}
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = state
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.get_state("light.test")
|
||||
|
||||
assert result == state
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_state_returns_none_for_404(self):
|
||||
"""get_state should return None for non-existent entity."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 404
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.get_state("light.nonexistent")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestHomeAssistantClientServices:
|
||||
"""Test service call methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_service_posts_to_correct_endpoint(self):
|
||||
"""call_service should POST to /api/services/{domain}/{service}."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.call_service("light", "turn_on", "light.test")
|
||||
|
||||
# Verify the correct URL was called
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/light/turn_on" in call_args[0][0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_turn_on_calls_correct_service(self):
|
||||
"""turn_on should call the turn_on service with attributes."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.turn_on("light.test", brightness=128)
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/light/turn_on" in call_args[0][0]
|
||||
# Check that brightness was passed in the payload
|
||||
payload = call_args[1]["json"]
|
||||
assert payload["entity_id"] == "light.test"
|
||||
assert payload["brightness"] == 128
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_turn_off_calls_correct_service(self):
|
||||
"""turn_off should call the turn_off service."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.turn_off("switch.test")
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/switch/turn_off" in call_args[0][0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_toggle_calls_correct_service(self):
|
||||
"""toggle should call the toggle service."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.toggle("light.test")
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/light/toggle" in call_args[0][0]
|
||||
|
||||
|
||||
class TestHomeAssistantClientScenes:
|
||||
"""Test scene methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_activate_scene_calls_scene_turn_on(self):
|
||||
"""activate_scene should call scene.turn_on service."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.activate_scene("scene.movie_night")
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/scene/turn_on" in call_args[0][0]
|
||||
|
||||
|
||||
class TestHomeAssistantClientScripts:
|
||||
"""Test script methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_script_calls_script_turn_on(self):
|
||||
"""run_script should call script.turn_on service."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.run_script("script.bedtime", {"delay": 5})
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/script/turn_on" in call_args[0][0]
|
||||
|
||||
|
||||
class TestHomeAssistantClientAutomations:
|
||||
"""Test automation methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enable_automation_calls_turn_on(self):
|
||||
"""enable_automation should call automation.turn_on."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.enable_automation("automation.motion")
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/automation/turn_on" in call_args[0][0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_automation_calls_turn_off(self):
|
||||
"""disable_automation should call automation.turn_off."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = []
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.disable_automation("automation.motion")
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/services/automation/turn_off" in call_args[0][0]
|
||||
|
||||
|
||||
class TestHomeAssistantClientHistory:
|
||||
"""Test history methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_history_calls_correct_endpoint(self):
|
||||
"""get_history should call /api/history/period endpoint."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = [[]]
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
await client.get_history("light.test", hours=24)
|
||||
|
||||
call_args = mock_client.get.call_args
|
||||
assert "/api/history/period/" in call_args[0][0]
|
||||
assert call_args[1]["params"]["filter_entity_id"] == "light.test"
|
||||
|
||||
|
||||
class TestHomeAssistantClientAreas:
|
||||
"""Test areas method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_areas_uses_template_api(self):
|
||||
"""get_areas should use the template API."""
|
||||
client = HomeAssistantClient(base_url="http://ha:8123", token="token")
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.text = '[{"id": "living_room", "name": "Living Room"}]'
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
mock_client.__aexit__.return_value = None
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
result = await client.get_areas()
|
||||
|
||||
call_args = mock_client.post.call_args
|
||||
assert "/api/template" in call_args[0][0]
|
||||
assert result == [{"id": "living_room", "name": "Living Room"}]
|
||||
|
||||
|
||||
class TestGetHomeAssistantClientSingleton:
|
||||
"""Test singleton pattern."""
|
||||
|
||||
def test_returns_same_instance(self):
|
||||
"""get_homeassistant_client should return singleton."""
|
||||
# Reset singleton
|
||||
import src.shared.clients.homeassistant_client as module
|
||||
module._homeassistant_client = None
|
||||
|
||||
client1 = get_homeassistant_client()
|
||||
client2 = get_homeassistant_client()
|
||||
|
||||
assert client1 is client2
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user