Compare commits
421 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 22474cdba5 | |||
| 0c839e8c6f | |||
| c63cf7db50 | |||
| d9d2b91384 | |||
| d6ea5fb1fd | |||
| 1883825b18 | |||
| bc71fd6fac | |||
| 28aa9ccf5a | |||
| 44dd80cb58 | |||
| 23485833d8 | |||
| e23dcdf4e1 | |||
| f05ac55875 | |||
| bcefeb4163 | |||
| 33d08faf70 | |||
| 8a58c61278 | |||
| 8231e750d9 | |||
| 6ce645d210 | |||
| 89ca9f10c7 | |||
| baabd1fa62 | |||
| f14fc37e75 | |||
| e07938098a | |||
| 943b9db5c7 | |||
| a4604d6a9a | |||
| 18204628cc | |||
| 93b415c53e | |||
| e7adfb462b | |||
| ed1d6528c6 | |||
| c4be7163d6 | |||
| 13f55fff47 | |||
| 0ec20b9c23 | |||
| 4c11163bff | |||
| adda76a2ff | |||
| 47962ed476 | |||
| 6a0c9bd669 | |||
| 76fbf0a755 | |||
| 187193fa6e | |||
| cd9c9539a2 | |||
| 555517c144 | |||
| fc1554140f | |||
| bc5e80c954 | |||
| ab79080f0b | |||
| 4c216dd1ca | |||
| a37a3122f9 | |||
| a905cf729e | |||
| 3a16775188 | |||
| b363d89768 | |||
| 27c39f9cfc | |||
| e8d5b16acc | |||
| 437ad840ef | |||
| c2a232d8f0 | |||
| adaedb70ef | |||
| 5178cf9cbf | |||
| 01a0ef46c9 | |||
| a4c429d53a | |||
| 1fc244e818 | |||
| 8fb4b67372 | |||
| fbd41e3eb4 | |||
| 9c57a94e9f | |||
| 84b7b64ec0 | |||
| 29ed0f2a3b | |||
| 6c32e5266c | |||
| 9e88acaa36 | |||
| 0d10caf489 | |||
| 245d79569e | |||
| 8e86cd255c | |||
| 6ee667c384 | |||
| 21285498ae | |||
| 9751b65dce | |||
| 612217ad89 | |||
| b9ea806c0d | |||
| 22f736ce20 | |||
| deb22bec0f | |||
| 9320e175d1 | |||
| d8f220d825 | |||
| 77b7c82563 | |||
| 3fc392a314 | |||
| edd8882fa0 | |||
| 760369b102 | |||
| 8da527f964 | |||
| 76fbd74d20 | |||
| 384beeca8a | |||
| 869efda214 | |||
| 51b5d723ac | |||
| 33b482ce91 | |||
| 0cb2eefd29 | |||
| 5f5dc9c851 | |||
| 4864d269e8 | |||
| 29e4bed9e6 | |||
| c8da0ab6c4 | |||
| b04a458975 | |||
| 2299bd51ba | |||
| 45fb0c753c | |||
| c82e628e6f | |||
| 383a874bcf | |||
| 18634387c7 | |||
| 4b09f611d6 | |||
| 5c998f5bf9 | |||
| c595a513d5 | |||
| 08bd8bf7f9 | |||
| 1b7308d091 | |||
| 1060cb60ed | |||
| b58696bb7e | |||
| 4d0834069b | |||
| fc41b51bf0 | |||
| 3dbb6321fc | |||
| 7282b91d99 | |||
| 9ad11a021c | |||
| 9f29ac15da | |||
| 3092038e40 | |||
| 738e01bb7c | |||
| 4d0c70de98 | |||
| 701bd57293 | |||
| 85f04447ea | |||
| d931f3071d | |||
| 4a356fe88e | |||
| ca22e9c9d2 | |||
| 4de312c170 | |||
| e7a89a853f | |||
| 7b23618ae7 | |||
| 8b4e1a7428 | |||
| f37813a317 | |||
| 6e600fdbcd | |||
| f4211ad452 | |||
| 058c501e4a | |||
| 74033243c9 | |||
| 9910fd4445 | |||
| 4a66a4a384 | |||
| fa20d00d14 | |||
| f59274ae64 | |||
| b89fb608b6 | |||
| 116cd22ff8 | |||
| 0094ba01cd | |||
| f6fb984ec6 | |||
| 0dba13a354 | |||
| 5500552993 | |||
| fed49dba6d | |||
| fa17b13413 | |||
| 38185b9659 | |||
| 07e7c6ea0f | |||
| b638dccd36 | |||
| 8f7e19fdb1 | |||
| ba76f09a1c | |||
| 8b08c3886c | |||
| 5577e19782 | |||
| 268651bab0 | |||
| fc873e2d6b | |||
| 7389344b6d | |||
| 0a8f1419a6 | |||
| 09938bede4 | |||
| 073013bc61 | |||
| 865b9411da | |||
| 48d1a7d05c | |||
| ab1d3a6aa1 | |||
| d117047711 | |||
| d1c187ab16 | |||
| ecd3ba5918 | |||
| a919ff8611 | |||
| 39cf01c3c9 | |||
| c49bb028c4 | |||
| 0baf7f7750 | |||
| 6e496e102d | |||
| 7cb88a3163 | |||
| ce1a73abce | |||
| 5fddc65468 | |||
| f065e2b8c0 | |||
| f4cf286bb5 | |||
| 549d13f469 | |||
| 45cd192188 | |||
| 030a39dd5a | |||
| ed15d53493 | |||
| ffb7ef0d21 | |||
| e265c86997 | |||
| 23c876b558 | |||
| 74c65068c8 | |||
| 48277369f2 | |||
| 5f39267781 | |||
| caf11cdb8a | |||
| 2f44306089 | |||
| 6173d42ddf | |||
| 9afd559394 | |||
| 25da0c149a | |||
| 457f0d29ae | |||
| dfbd3e60a8 | |||
| 075eb6a76b | |||
| 4571bebf8e | |||
| ec45283257 | |||
| cce3fa773a | |||
| 50474f7b13 | |||
| 9c7043bab1 | |||
| 2dd2f2ab06 | |||
| e555561a2d | |||
| 4ac3d593aa | |||
| 1a3860a4d8 | |||
| 3d9ff44d1a | |||
| 12378def4d | |||
| f6f7853aa4 | |||
| 802d8f1e8c | |||
| 4f9aa7e3c2 | |||
| cb25b94cb7 | |||
| 01aaf4c78f | |||
| 96c8dd7402 | |||
| f5c2c95af0 | |||
| 06fe8623bc | |||
| f2fed518f0 | |||
| cc2a638c76 | |||
| 5cec4a7a6f | |||
| 7e0df57f8c | |||
| aea1ff95f6 | |||
| c8fdca7f60 | |||
| 014b88ee56 | |||
| b6bda3d692 | |||
| d7fb51f427 | |||
| 0d57e3501a | |||
| 3672312028 | |||
| c5f117e5b1 | |||
| 10a5c29702 | |||
| 312a646b89 | |||
| 01adc9a00f | |||
| 2e9ca52cdb | |||
| cd4eba9803 | |||
| 18646e3d1b | |||
| 8c5e1b931e | |||
| d1be2e4951 | |||
| 953ea05756 | |||
| 507b71c586 | |||
| aebcf25bf4 | |||
| ae42cac61e | |||
| e2ad7d7fb6 | |||
| 9e1334eb6d | |||
| c41993310b | |||
| cb25f21c44 | |||
| 392e85ead4 | |||
| a4bf8afac9 | |||
| a559470369 | |||
| 8cab17472e | |||
| e9d404b1ff | |||
| dab6c74046 | |||
| ca8b255148 | |||
| 02a2ad6df5 | |||
| 4ef0f108ea | |||
| dc8ef0e463 | |||
| 6a7657aeda | |||
| eca8b8815b | |||
| d0f7a97f92 | |||
| 765cb965e6 | |||
| 063a839790 | |||
| ae41a64e66 | |||
| 0901b1e832 | |||
| 7cc720786e | |||
| e167a6be12 | |||
| 0fa926284c | |||
| 5c17de0c3c | |||
| 8efadc4432 | |||
| 1e40540ef4 | |||
| 7cbbb41661 | |||
| 952a9f3234 | |||
| ab8872f79e | |||
| be4893e2a7 | |||
| b3c6a5fdc9 | |||
| b7d17cea78 | |||
| 36d6448f5f | |||
| 20a5f6a9a1 | |||
| 1c94583307 | |||
| 95a7454bee | |||
| 649496b762 | |||
| d13e16f5e1 | |||
| d5f9df33b7 | |||
| 468e0eacda | |||
| 99097090e6 | |||
| 2a9e57ad0d | |||
| e4c5e7f2db | |||
| 70957e462a | |||
| 6f35eb77ae | |||
| 684a11610a | |||
| 575e5837ea | |||
| 996ea73bbf | |||
| d0e5feeaa5 | |||
| 1f784b552d | |||
| 10cbe63095 | |||
| 4547105f3b | |||
| dacf105200 | |||
| 2525c58471 | |||
| 8dd350286e | |||
| 5af4de0d7e | |||
| ebfe991a15 | |||
| cc1507a33e | |||
| ae377baa74 | |||
| a0d4e76662 | |||
| b61c2256a8 | |||
| 7e6588a99d | |||
| 56b54e269a | |||
| 6e597b9e21 | |||
| a2c3d6877e | |||
| a200955ef2 | |||
| 76266fc3d0 | |||
| 9d54a2542b | |||
| 520078fd64 | |||
| 55c4dc1281 | |||
| 74ac7b00bb | |||
| ecc0acc8dd | |||
| fba01ddfb2 | |||
| 35a251a0b4 | |||
| b10eadf64b | |||
| 50fcf5c077 | |||
| 19db7db8d0 | |||
| 189f29ee41 | |||
| f2c3264be6 | |||
| f07a632c86 | |||
| 0b68efb6e0 | |||
| 54ec87a836 | |||
| 3d64ec9061 | |||
| ac6bd3304d | |||
| 6ec35988cc | |||
| e985f0122e | |||
| 2bd778117f | |||
| 23ae12e69c | |||
| 65d4fad3c5 | |||
| e6f64c39f3 | |||
| e7c42c17b9 | |||
| f7be50952a | |||
| 402e662c0c | |||
| 7450dd0ce5 | |||
| 7a88639250 | |||
| 1d33735e2f | |||
| d11b43b69f | |||
| 294d02f9fb | |||
| a10029f36c | |||
| 6a2ebae8c9 | |||
| 1a922d4171 | |||
| d906c12aa9 | |||
| 4906ce0cd9 | |||
| e7804c0f58 | |||
| b40eb3e88c | |||
| 98795e31dd | |||
| b4a7627718 | |||
| 9d27cd9fe8 | |||
| e3ef852eca | |||
| f824a1e6fe | |||
| ad09ffa6ec | |||
| 63ae706dd0 | |||
| 74b5d0dc8c | |||
| 25662e525c | |||
| d4e992a9e2 | |||
| 35ada0e662 | |||
| c3e2264771 | |||
| 76741d3ee6 | |||
| 62752a8390 | |||
| 461cb30b28 | |||
| 94aa88c154 | |||
| 4f695d7e62 | |||
| 7a48180dc3 | |||
| a0b0944709 | |||
| b79da51269 | |||
| 40a940304b | |||
| e344e961d6 | |||
| d6b3e8b804 | |||
| f997b3f7c5 | |||
| c795f8f873 | |||
| c52367401b | |||
| e9e0e87013 | |||
| 2b5331a1a0 | |||
| 6f41fa7cbe | |||
| cccc4a9d5a | |||
| 72b4a5e2bd | |||
| 8f1d2a149a | |||
| 7261c75bb2 | |||
| 84bf7e4aeb | |||
| 9373d93169 | |||
| 7fc8b82621 | |||
| 2d078c8b1e | |||
| 92c8dfe986 | |||
| 0bc0d99c21 | |||
| 965f6f6585 | |||
| a9f4657d03 | |||
| 875594d66d | |||
| 762e0de44c | |||
| 955dadc604 | |||
| 83f94b1f09 | |||
| 6807f449b7 | |||
| 92d0d5b891 | |||
| 834982d423 | |||
| 9873a8186a | |||
| d4b52668aa | |||
| d0191cd549 | |||
| 0ae0e3fec1 | |||
| 0926e4de83 | |||
| 8b70daed53 | |||
| ac6c97b6ce | |||
| a433c82488 | |||
| f8700fd7ed | |||
| 0fcfc745ff | |||
| d214ab82db | |||
| 753f1506b6 | |||
| 58bf30ed15 | |||
| c067c03662 | |||
| d273535950 | |||
| 7f97ba8e9b | |||
| caf73e39ba | |||
| 1e70462c2b | |||
| b9684b0107 | |||
| 716b8fa631 | |||
| d724a92d34 | |||
| 2ce7862058 | |||
| 285d3dace8 | |||
| 9fefe289a7 | |||
| 843683d579 | |||
| c1a4d2d9af | |||
| ea6c466c6c | |||
| 29c3563148 | |||
| 0509b9eb4a | |||
| 137757602f | |||
| 3cd8674c31 | |||
| 899fba9c9b | |||
| a2c787d474 | |||
| 8ba2b48967 | |||
| ec9b73225c | |||
| 0296ea5630 | |||
| 6dbd55a9ac | |||
| 27f4ecce86 | |||
| 7b72ccdc3c | |||
| 6b302b3279 |
@@ -0,0 +1,195 @@
|
||||
---
|
||||
name: sift-backlog
|
||||
description: Triage and organize backlog tasks into actionable plans. Use when asked to review the backlog, prioritize tasks, create plans from backlog items, or move tasks from backlog to open status. Handles the full workflow of listing backlog tasks, grouping related tasks into plans, setting priorities and dependencies, activating plans, and changing task status from backlog to open.
|
||||
---
|
||||
|
||||
# Sift Backlog
|
||||
|
||||
Triage backlog tasks: prioritize, group into plans, set dependencies, and activate.
|
||||
|
||||
## Overview
|
||||
|
||||
1. List backlog tasks (`sf task backlog`)
|
||||
2. Clarify and enrich each task (titles, descriptions)
|
||||
3. Identify groupings and create draft plans
|
||||
4. Add tasks to plans and set dependencies
|
||||
5. Activate plans
|
||||
6. Set task status to open
|
||||
|
||||
## Workflow
|
||||
|
||||
### Step 1: List Backlog Tasks
|
||||
|
||||
```bash
|
||||
sf task backlog
|
||||
```
|
||||
|
||||
### Step 2: Clarify and Enrich Tasks
|
||||
|
||||
Backlog tasks often have only a brief title with no description. Before organizing, ensure each task is well-defined.
|
||||
|
||||
**For each task, evaluate:**
|
||||
|
||||
- Is the title clear and actionable?
|
||||
- Is there a description? Check with `sf task describe <task-id> --show`
|
||||
- Is the scope unambiguous?
|
||||
|
||||
**If the title is unclear**, update it:
|
||||
|
||||
```bash
|
||||
sf update <task-id> --title "Clear, actionable title"
|
||||
```
|
||||
|
||||
**Add a description** with context, scope, and acceptance criteria:
|
||||
|
||||
```bash
|
||||
sf task describe <task-id> --content "Description with:
|
||||
- What needs to be done
|
||||
- Why it matters
|
||||
- Acceptance criteria
|
||||
- Any relevant context"
|
||||
```
|
||||
|
||||
**Use your best judgment** to interpret tasks and make reasonable decisions about scope, grouping, and priority. You have context about the codebase, project patterns, and typical development practices—leverage this knowledge rather than deferring to the user for routine decisions.
|
||||
|
||||
**Only ask the user for clarity when absolutely necessary:**
|
||||
|
||||
- The task is fundamentally ambiguous (multiple mutually exclusive interpretations)
|
||||
- Critical business logic or user-facing behavior that could go wrong in meaningful ways
|
||||
- External dependencies or integrations you cannot verify
|
||||
|
||||
**Do NOT ask about:**
|
||||
|
||||
- Implementation details you can reasonably infer
|
||||
- Priority or grouping decisions—use your judgment
|
||||
- Standard development practices (testing, code style, etc.)
|
||||
- Tasks where a reasonable interpretation exists
|
||||
|
||||
### Step 3: Create Draft Plans
|
||||
|
||||
Group related tasks into plans using your best judgment. Plans start as drafts (tasks won't be dispatched until activated).
|
||||
|
||||
**Grouping guidance:**
|
||||
|
||||
- Group tasks that share a common theme, feature area, or goal
|
||||
- Consider technical dependencies when grouping (tasks that touch the same files/modules)
|
||||
- Separate unrelated work into distinct plans for parallel execution
|
||||
- Don't over-group—if tasks are truly independent, separate plans enable better parallelism
|
||||
- Don't under-group—related tasks benefit from shared context and coordinated execution
|
||||
|
||||
```bash
|
||||
sf plan create --title "Plan Name"
|
||||
```
|
||||
|
||||
**Example:**
|
||||
|
||||
```bash
|
||||
sf plan create --title "Authentication Improvements"
|
||||
# Output: Created plan el-abc123
|
||||
```
|
||||
|
||||
### Step 4: Add Tasks to Plans
|
||||
|
||||
```bash
|
||||
sf plan add-task <plan-id> <task-id>
|
||||
```
|
||||
|
||||
**Example:**
|
||||
|
||||
```bash
|
||||
sf plan add-task el-abc123 el-task1
|
||||
sf plan add-task el-abc123 el-task2
|
||||
```
|
||||
|
||||
### Step 5: Set Dependencies Between Tasks
|
||||
|
||||
Use `blocks` dependency when one task must complete before another can start.
|
||||
|
||||
```bash
|
||||
sf dependency add <blocked-id> <blocker-id> --type blocks
|
||||
```
|
||||
|
||||
**Semantics:** The first ID is blocked BY the second ID. The blocker must complete first.
|
||||
|
||||
**Example:** Task 2 can't start until Task 1 completes:
|
||||
|
||||
```bash
|
||||
sf dependency add el-task2 el-task1 --type blocks
|
||||
```
|
||||
|
||||
### Step 6: Update Priorities
|
||||
|
||||
Set priorities based on your assessment of impact, urgency, and dependencies. Use your judgment—you don't need user confirmation for routine prioritization.
|
||||
|
||||
**Priority guidance:**
|
||||
|
||||
- **Critical (1):** Blocking issues, security vulnerabilities, production bugs
|
||||
- **High (2):** Important features with deadlines, significant user impact
|
||||
- **Medium (3):** Standard feature work, most tasks default here
|
||||
- **Low (4):** Nice-to-haves, minor improvements, tech debt
|
||||
- **Minimal (5):** Backlog cleanup, documentation, exploratory work
|
||||
|
||||
```bash
|
||||
sf update <task-id> --priority <1-5>
|
||||
```
|
||||
|
||||
| Value | Level |
|
||||
| ----- | -------- |
|
||||
| 1 | Critical |
|
||||
| 2 | High |
|
||||
| 3 | Medium |
|
||||
| 4 | Low |
|
||||
| 5 | Minimal |
|
||||
|
||||
### Step 7: Activate Plans
|
||||
|
||||
Once tasks are organized with dependencies set, activate plans to enable dispatch.
|
||||
|
||||
```bash
|
||||
sf plan activate <plan-id>
|
||||
```
|
||||
|
||||
### Step 8: Set Task Status to Open
|
||||
|
||||
Move tasks from backlog to open so they become ready for work.
|
||||
|
||||
```bash
|
||||
sf update <id> --status open
|
||||
```
|
||||
|
||||
## Other Actions
|
||||
|
||||
**Close obsolete tasks:**
|
||||
|
||||
```bash
|
||||
sf task close <id> --reason "Won't do: <reason>"
|
||||
```
|
||||
|
||||
**Defer tasks:**
|
||||
|
||||
```bash
|
||||
sf task defer <id> --until <date>
|
||||
```
|
||||
|
||||
**View existing plans:**
|
||||
|
||||
```bash
|
||||
sf plan list
|
||||
```
|
||||
|
||||
**View tasks in a plan:**
|
||||
|
||||
```bash
|
||||
sf plan tasks <plan-id>
|
||||
```
|
||||
|
||||
## Tips
|
||||
|
||||
- **Use your best judgment** for grouping, prioritization, and task interpretation—don't defer routine decisions to the user
|
||||
- **Only escalate to the user** when ambiguity is fundamental and could lead to wasted work (mutually exclusive interpretations, critical business decisions)
|
||||
- Make reasonable inferences about implementation details, scope, and priority based on codebase context
|
||||
- Create plans before setting dependencies to avoid dispatch race conditions
|
||||
- Always activate plans after dependencies are set
|
||||
- Focus on oldest backlog items first (sorted by creation date)
|
||||
- Every task should have a clear title and description before activation
|
||||
- When uncertain about a minor detail, make a reasonable choice and document it in the task description—workers can ask if needed
|
||||
+8
-7
@@ -6,10 +6,9 @@ POSTGRES_DB=headquarter
|
||||
# Redis Configuration
|
||||
REDIS_URL=redis://redis:6379/0
|
||||
|
||||
# JWT Configuration
|
||||
JWT_SECRET=change-me-in-production
|
||||
JWT_ALGORITHM=HS256
|
||||
JWT_EXPIRATION_HOURS=24
|
||||
# Session Configuration
|
||||
SESSION_SECRET=change-me-in-production
|
||||
SESSION_TTL_HOURS=24
|
||||
|
||||
# Application Configuration
|
||||
APP_ENV=development
|
||||
@@ -27,14 +26,16 @@ AUTHENTIK_DOMAIN=authentik.local
|
||||
# WEB_PUBLIC_URL=https://app.example.com
|
||||
|
||||
# Authentik Configuration
|
||||
# Client ID: The OAuth client ID from Authentik (may be a UUID)
|
||||
AUTHENTIK_CLIENT_ID=headquarter-web
|
||||
AUTHENTIK_CLIENT_SECRET=change-me
|
||||
# Application Slug: The URL-friendly identifier used in Authentik URLs
|
||||
# This is often the same as the application identifier/slug in Authentik
|
||||
# e.g., if your Authentik app URL is /application/o/headquarter-web/, use "headquarter-web"
|
||||
AUTHENTIK_APPLICATION_SLUG=headquarter-web
|
||||
# Override Authentik URLs if they differ from the default pattern
|
||||
# AUTHENTIK_AUTHORIZE_URL=https://authentik.example.com/application/o/authorize/
|
||||
# AUTHENTIK_TOKEN_URL=https://authentik.example.com/application/o/token/
|
||||
# AUTHENTIK_JWKS_URL=https://authentik.example.com/application/o/headquarter-web/jwks/
|
||||
# AUTHENTIK_ISSUER=https://authentik.example.com/application/o/headquarter-web/
|
||||
AUTHENTIK_AUDIENCE=headquarter-web
|
||||
|
||||
# Frontend Configuration
|
||||
VITE_API_BASE_URL=http://localhost:8000
|
||||
|
||||
@@ -48,3 +48,4 @@ apps/web/dist/
|
||||
# OS
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
/.stoneforge/.worktrees/
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
262629
|
||||
1779624255076
|
||||
@@ -0,0 +1,6 @@
|
||||
# Runtime data
|
||||
*.db
|
||||
*.db-journal
|
||||
*.db-wal
|
||||
*.db-shm
|
||||
daemon-state.json
|
||||
@@ -0,0 +1,20 @@
|
||||
# Stoneforge Configuration
|
||||
|
||||
database: stoneforge.db
|
||||
sync:
|
||||
auto_export: true
|
||||
elements_file: elements.jsonl
|
||||
dependencies_file: dependencies.jsonl
|
||||
playbooks:
|
||||
paths:
|
||||
- playbooks
|
||||
identity:
|
||||
mode: soft
|
||||
merge:
|
||||
auto_merge: true
|
||||
target_branch: null
|
||||
require_approval: false
|
||||
workflow:
|
||||
preset: auto
|
||||
agents:
|
||||
permission_model: unrestricted
|
||||
@@ -0,0 +1,43 @@
|
||||
{"blockedId":"el-1of","blockerId":"el-258","type":"parent-child","createdAt":"2026-05-24T09:44:58.759Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5fe","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:40.892Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1nj","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.010Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1bn","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.127Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4hr","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.244Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-62c","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.372Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5z8","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.490Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1t7","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.607Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5j5","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.726Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-2xl","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.844Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4bc","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:41.959Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-107","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:42.074Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-32e","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:42.195Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-3ou","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:42.311Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-14w","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:42.425Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1ou","blockerId":"el-20no","type":"parent-child","createdAt":"2026-05-24T12:44:42.541Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1nj","blockerId":"el-5fe","type":"blocks","createdAt":"2026-05-24T12:44:42.651Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1bn","blockerId":"el-5fe","type":"blocks","createdAt":"2026-05-24T12:44:42.761Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4hr","blockerId":"el-5fe","type":"blocks","createdAt":"2026-05-24T12:44:42.868Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-62c","blockerId":"el-1nj","type":"blocks","createdAt":"2026-05-24T12:44:42.979Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-62c","blockerId":"el-1bn","type":"blocks","createdAt":"2026-05-24T12:44:43.092Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5z8","blockerId":"el-1nj","type":"blocks","createdAt":"2026-05-24T12:44:43.205Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5z8","blockerId":"el-4hr","type":"blocks","createdAt":"2026-05-24T12:44:43.313Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1t7","blockerId":"el-1bn","type":"blocks","createdAt":"2026-05-24T12:44:43.422Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1t7","blockerId":"el-4hr","type":"blocks","createdAt":"2026-05-24T12:44:43.529Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1t7","blockerId":"el-62c","type":"blocks","createdAt":"2026-05-24T12:44:43.647Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-5j5","blockerId":"el-1t7","type":"blocks","createdAt":"2026-05-24T12:44:43.758Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-2xl","blockerId":"el-1t7","type":"blocks","createdAt":"2026-05-24T12:44:43.876Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4bc","blockerId":"el-1nj","type":"blocks","createdAt":"2026-05-24T12:44:43.987Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4bc","blockerId":"el-1bn","type":"blocks","createdAt":"2026-05-24T12:44:44.096Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-4bc","blockerId":"el-62c","type":"blocks","createdAt":"2026-05-24T12:44:44.208Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-107","blockerId":"el-5z8","type":"blocks","createdAt":"2026-05-24T12:44:44.319Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-32e","blockerId":"el-5j5","type":"blocks","createdAt":"2026-05-24T12:44:44.429Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-32e","blockerId":"el-2xl","type":"blocks","createdAt":"2026-05-24T12:44:44.539Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-3ou","blockerId":"el-4bc","type":"blocks","createdAt":"2026-05-24T12:44:44.650Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-3ou","blockerId":"el-107","type":"blocks","createdAt":"2026-05-24T12:44:44.761Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-14w","blockerId":"el-32e","type":"blocks","createdAt":"2026-05-24T12:44:44.873Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1ou","blockerId":"el-3ou","type":"blocks","createdAt":"2026-05-24T12:44:44.987Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-1ou","blockerId":"el-14w","type":"blocks","createdAt":"2026-05-24T12:44:45.107Z","createdBy":"el-2jua"}
|
||||
{"blockedId":"el-375","blockerId":"el-26p","type":"replies-to","createdAt":"2026-05-24T13:21:42.486Z","createdBy":"el-2i1s"}
|
||||
{"blockedId":"el-3n4","blockerId":"el-31p","type":"replies-to","createdAt":"2026-05-24T13:21:46.044Z","createdBy":"el-13ju"}
|
||||
{"blockedId":"el-3jer","blockerId":"el-1xx","type":"replies-to","createdAt":"2026-05-24T13:24:47.580Z","createdBy":"el-4350"}
|
||||
{"blockedId":"el-1afv","blockerId":"el-1ozw","type":"replies-to","createdAt":"2026-05-24T13:32:42.658Z","createdBy":"el-51a8"}
|
||||
File diff suppressed because one or more lines are too long
@@ -87,6 +87,31 @@ Do not claim completion without verification evidence.
|
||||
|
||||
## Git workflow
|
||||
|
||||
### Branching strategy
|
||||
|
||||
For every spec change or new functionality:
|
||||
|
||||
1. Create a new branch from `dev` with a proper prefix:
|
||||
- `feat/` for new features (e.g., `feat/tool-workshop`)
|
||||
- `fix/` for bug fixes (e.g., `fix/terminal-tty`)
|
||||
- `refactor/` for refactors (e.g., `refactor/api-cleanup`)
|
||||
- `docs/` for documentation (e.g., `docs/api-guide`)
|
||||
- `chore/` for maintenance (e.g., `chore/update-deps`)
|
||||
2. Branch name should reference the OpenSpec change name when applicable.
|
||||
3. Do not commit directly to `main` or `dev`.
|
||||
|
||||
### Completion and merge
|
||||
|
||||
When implementation is complete and verified:
|
||||
|
||||
1. Ensure all tests pass and quality gates are met.
|
||||
2. Stage all changes with `git add -A`.
|
||||
3. Create a commit with a proper conventional commit message (see below).
|
||||
4. Switch to `dev`: `git checkout dev`.
|
||||
5. Merge the feature branch: `git merge --no-ff <branch-name>`.
|
||||
6. Push to remote: `git push origin dev`.
|
||||
7. Delete the local feature branch if desired: `git branch -d <branch-name>`.
|
||||
|
||||
### Auto-commit on spec completion
|
||||
|
||||
When an OpenSpec change is fully implemented and all tasks are complete:
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
# Changelog
|
||||
|
||||
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).
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### Added
|
||||
|
||||
- **Project Management** - Create and manage projects with dashboard view
|
||||
- **Git Repository Management** - Bare repository initialization and mirror cloning with smart URL parsing
|
||||
- **Repository Workspace** - File browser with syntax highlighting, branch switching, and file editing
|
||||
- **Git History Visualization** - Commit history with graph visualization and diff viewing
|
||||
- **Git Control** - Branch management, commit, fetch/pull/push, merge operations
|
||||
- **File Editor** - Syntax highlighting for 50+ languages with edit/commit workflow
|
||||
- **Smart Git URL Parsing** - Automatic detection and correction of browser URLs to git clone URLs
|
||||
- **OAuth2 Authentication** - Session-based authentication via Authentik with simplified flow
|
||||
- **User Profile** - Profile management with avatar upload
|
||||
- **User Settings** - Theme selection, git identity, and preference management
|
||||
- **SSH Key Management** - Ed25519 key generation with secure storage
|
||||
- **Tool Types** - Built-in development tools (code-server, jupyter-notebook) with custom type support
|
||||
- **Comprehensive Documentation** - Architecture, API, deployment, and development guides
|
||||
|
||||
### Changed
|
||||
|
||||
- Simplified authentication from JWT to session-based cookies
|
||||
- Restructured test infrastructure with unit/integration/system separation
|
||||
- Improved Docker deployment with Traefik integration
|
||||
|
||||
### Fixed
|
||||
|
||||
- Database migration chain errors
|
||||
- Cross-origin cookie handling for OAuth flow
|
||||
- Nginx permission issues in container
|
||||
|
||||
## [0.1.0] - 2026-05-19
|
||||
|
||||
### Added
|
||||
|
||||
- Initial release with core project and repository management
|
||||
- OAuth2 authentication with Authentik
|
||||
- Basic file browsing and git history viewing
|
||||
- Development tool type definitions
|
||||
@@ -1,53 +1,195 @@
|
||||
# Headquarter
|
||||
|
||||
## Testing Strategy
|
||||
A self-hosted platform for managing projects, git repositories, and development tools with OAuth2 authentication.
|
||||
|
||||
The project uses a three-tier testing approach:
|
||||
## Overview
|
||||
|
||||
### Test Categories
|
||||
Headquarter provides a centralized workspace for development teams to:
|
||||
- Manage projects and their associated git repositories
|
||||
- Browse repository files and view git history
|
||||
- Spawn development tools (VS Code Server, Jupyter Notebook, etc.)
|
||||
- Manage SSH keys and user preferences
|
||||
|
||||
1. **Unit Tests** (`apps/api/tests/unit/`)
|
||||
- Fast tests with no external dependencies
|
||||
- Use SQLite in-memory database
|
||||
- Run with: `make test-unit` or `pytest -m unit`
|
||||
## Features
|
||||
|
||||
2. **Integration Tests** (`apps/api/tests/integration/`)
|
||||
- Test API endpoints with database
|
||||
- Use PostgreSQL with transaction rollback
|
||||
- Run with: `make test-integration` or `pytest -m integration`
|
||||
### Project Management
|
||||
- Create and manage projects
|
||||
- View all projects in a dashboard
|
||||
- Click any project to open its workspace
|
||||
|
||||
3. **System/E2E Tests** (`e2e/`)
|
||||
- End-to-end tests using Playwright
|
||||
- Test full user journeys
|
||||
- Run with: `make test-e2e`
|
||||
### Git Repository Management
|
||||
- Initialize bare repositories
|
||||
- Clone repositories (including mirror clones)
|
||||
- Smart URL parsing (converts browser URLs to git URLs)
|
||||
- View repository history and commit details
|
||||
|
||||
### Running Tests
|
||||
### Repository Workspace
|
||||
- Browse files and directories
|
||||
- View file contents with syntax highlighting
|
||||
- Switch between branches
|
||||
- Quick file editing with automatic commits
|
||||
|
||||
```bash
|
||||
# Run all tests (excludes system tests by default)
|
||||
make test
|
||||
### Git History Visualization
|
||||
- View commit history with branch graph
|
||||
- See commit details, statistics, and diffs
|
||||
- Filter by branch
|
||||
|
||||
# Run specific categories
|
||||
make test-unit # Fast unit tests only
|
||||
make test-integration # Integration tests with DB
|
||||
make test-system # Full stack tests
|
||||
make test-e2e # Browser-based E2E tests
|
||||
### Authentication
|
||||
- OAuth2 via Authentik
|
||||
- Session-based authentication
|
||||
- User profile management
|
||||
|
||||
# Inside Docker container
|
||||
docker compose exec api pytest -v -m unit
|
||||
docker compose exec api pytest -v -m integration
|
||||
### Tool Management
|
||||
- Built-in tool types (code-server, jupyter-notebook)
|
||||
- Create custom tool types with Docker Compose templates
|
||||
- Template validation
|
||||
|
||||
### User Settings
|
||||
- Theme selection (system/light/dark)
|
||||
- Git identity configuration
|
||||
- Default editor preference
|
||||
|
||||
### SSH Key Management
|
||||
- Generate Ed25519 key pairs
|
||||
- Copy public keys to clipboard
|
||||
- Delete keys
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Prerequisites
|
||||
- Docker and Docker Compose
|
||||
- Git
|
||||
|
||||
### Local Development
|
||||
|
||||
1. **Clone the repository:**
|
||||
```bash
|
||||
git clone <repository-url>
|
||||
cd headquarter
|
||||
```
|
||||
|
||||
2. **Set up environment:**
|
||||
```bash
|
||||
cp .env.example .env
|
||||
# Edit .env with your settings
|
||||
```
|
||||
|
||||
3. **Start services:**
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
4. **Access the application:**
|
||||
- Frontend: http://localhost:5173
|
||||
- API: http://localhost:8000
|
||||
- API Docs: http://localhost:8000/docs
|
||||
|
||||
### Production Deployment
|
||||
|
||||
See [Deployment Guide](docs/deployment/) for production setup with Traefik and Authentik.
|
||||
|
||||
## Tech Stack
|
||||
|
||||
### Backend
|
||||
- **FastAPI** - Python web framework
|
||||
- **SQLAlchemy** - ORM with async PostgreSQL support
|
||||
- **Pydantic** - Data validation
|
||||
- **Alembic** - Database migrations
|
||||
- **python-jose** - JWT handling
|
||||
|
||||
### Frontend
|
||||
- **React** - UI library
|
||||
- **TypeScript** - Type safety
|
||||
- **Vite** - Build tool
|
||||
- **React Router** - Client-side routing
|
||||
|
||||
### Infrastructure
|
||||
- **Docker** - Containerization
|
||||
- **PostgreSQL** - Database
|
||||
- **Traefik** - Reverse proxy (production)
|
||||
- **Authentik** - Identity provider
|
||||
|
||||
## Documentation
|
||||
|
||||
- [User Guide](docs/features/) - Feature documentation
|
||||
- [API Reference](docs/api/) - API endpoints
|
||||
- [Architecture](docs/architecture/) - System design
|
||||
- [Deployment](docs/deployment/) - Setup guides
|
||||
- [Development](docs/development/) - Contributing
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
.
|
||||
├── apps/
|
||||
│ ├── api/ # FastAPI backend
|
||||
│ │ ├── src/
|
||||
│ │ │ ├── api/ # API routes
|
||||
│ │ │ ├── auth/ # Authentication
|
||||
│ │ │ ├── models/ # Database models
|
||||
│ │ │ └── utils/ # Utilities
|
||||
│ │ ├── tests/ # Test suite
|
||||
│ │ └── Dockerfile
|
||||
│ └── web/ # React frontend
|
||||
│ ├── src/
|
||||
│ │ ├── api/ # API clients
|
||||
│ │ ├── components/# UI components
|
||||
│ │ └── pages/ # Page components
|
||||
│ └── Dockerfile
|
||||
├── docs/ # Documentation
|
||||
├── docker-compose.yml # Development setup
|
||||
├── docker-compose.traefik.yml # Production setup
|
||||
└── Makefile # Common commands
|
||||
```
|
||||
|
||||
### Test Markers
|
||||
## Development
|
||||
|
||||
Tests are marked with pytest markers:
|
||||
- `@pytest.mark.unit` - Fast, isolated tests
|
||||
- `@pytest.mark.integration` - Tests with database/external services
|
||||
- `@pytest.mark.system` - Full stack tests
|
||||
### Backend Development
|
||||
```bash
|
||||
cd apps/api
|
||||
python -m venv .venv
|
||||
source .venv/bin/activate
|
||||
pip install -e ".[dev]"
|
||||
uvicorn src.main:app --reload
|
||||
```
|
||||
|
||||
### Shared Fixtures
|
||||
### Frontend Development
|
||||
```bash
|
||||
cd apps/web
|
||||
npm install
|
||||
npm run dev
|
||||
```
|
||||
|
||||
Common fixtures are in `apps/api/tests/conftest.py`:
|
||||
- `sqlite_engine` - SQLite engine for unit tests
|
||||
- `postgres_engine` - PostgreSQL engine for integration tests
|
||||
- `db_session` - Database session with transaction rollback
|
||||
- `test_client` - FastAPI TestClient instance
|
||||
### Running Tests
|
||||
```bash
|
||||
# Backend tests
|
||||
make test
|
||||
|
||||
# Frontend tests
|
||||
make test-web
|
||||
|
||||
# All quality gates
|
||||
make lint
|
||||
make typecheck
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Key environment variables:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `API_DOMAIN` | API domain | `localhost` |
|
||||
| `WEB_DOMAIN` | Web domain | `localhost` |
|
||||
| `AUTHENTIK_DOMAIN` | Authentik domain | - |
|
||||
| `AUTHENTIK_CLIENT_ID` | OAuth client ID | - |
|
||||
| `AUTHENTIK_CLIENT_SECRET` | OAuth client secret | - |
|
||||
| `DATABASE_URL` | PostgreSQL URL | - |
|
||||
| `JWT_SECRET` | JWT signing secret | - |
|
||||
| `REPO_BASE_PATH` | Repository storage path | `/data/repos` |
|
||||
|
||||
See [Environment Variables](docs/deployment/environment.md) for complete list.
|
||||
|
||||
## License
|
||||
|
||||
[License information]
|
||||
|
||||
+34
-11
@@ -16,29 +16,51 @@ RUN pip install --no-cache-dir --user -e ".[dev]"
|
||||
# Production stage
|
||||
FROM python:3.11-slim
|
||||
|
||||
# Create non-root user
|
||||
RUN groupadd -r appgroup && useradd -r -g appgroup appuser
|
||||
# Create non-root user and add to docker group
|
||||
RUN groupadd -r appgroup && useradd -r -g appgroup appuser \
|
||||
&& groupadd -r docker || true \
|
||||
&& usermod -aG docker appuser
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install runtime dependencies
|
||||
# Install runtime dependencies including Docker CLI
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
libpq5 \
|
||||
git \
|
||||
openssh-client \
|
||||
netcat-openbsd \
|
||||
ca-certificates \
|
||||
curl \
|
||||
gnupg \
|
||||
&& install -m 0755 -d /etc/apt/keyrings \
|
||||
&& curl -fsSL https://download.docker.com/linux/debian/gpg | gpg --dearmor -o /etc/apt/keyrings/docker.gpg \
|
||||
&& chmod a+r /etc/apt/keyrings/docker.gpg \
|
||||
&& echo "deb [arch="$(dpkg --print-architecture)" signed-by=/etc/apt/keyrings/docker.gpg] https://download.docker.com/linux/debian \
|
||||
"$(. /etc/os-release && echo "$VERSION_CODENAME")" stable" > /etc/apt/sources.list.d/docker.list \
|
||||
&& apt-get update \
|
||||
&& apt-get install -y --no-install-recommends docker-ce-cli docker-compose-plugin \
|
||||
&& curl -L --output /usr/local/bin/cloudflared https://github.com/cloudflare/cloudflared/releases/latest/download/cloudflared-linux-amd64 \
|
||||
&& chmod +x /usr/local/bin/cloudflared \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Copy dependencies from builder
|
||||
COPY --from=builder /root/.local /home/appuser/.local
|
||||
ENV PATH=/home/appuser/.local/bin:$PATH
|
||||
COPY --from=builder /root/.local /root/.local
|
||||
ENV PATH=/root/.local/bin:$PATH
|
||||
|
||||
# Copy application code
|
||||
COPY --chown=appuser:appgroup . .
|
||||
|
||||
# Create directories for repo storage
|
||||
RUN mkdir -p /data/repos && chown -R appuser:appgroup /data/repos
|
||||
# Create directories for repo and instance storage
|
||||
RUN mkdir -p /data/repos /data/instances && chown -R appuser:appgroup /data
|
||||
|
||||
# Switch to non-root user
|
||||
USER appuser
|
||||
# Copy wait-for-db script
|
||||
COPY wait-for-db.sh /usr/local/bin/wait-for-db.sh
|
||||
RUN chmod +x /usr/local/bin/wait-for-db.sh
|
||||
|
||||
# NOTE: Running as root to access Docker socket for managing tool instances
|
||||
# This is required because Docker socket permissions require root or docker group membership
|
||||
# which doesn't work well across container boundaries.
|
||||
# Consider using Docker-in-Docker or rootless Docker for production hardening.
|
||||
|
||||
# Expose port
|
||||
EXPOSE 8000
|
||||
@@ -47,5 +69,6 @@ EXPOSE 8000
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health')" || exit 1
|
||||
|
||||
# Run the application
|
||||
CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||
# Run the application (with database wait)
|
||||
ENTRYPOINT ["/usr/local/bin/wait-for-db.sh"]
|
||||
CMD ["uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
# Headquarter API
|
||||
|
||||
The backend API for Headquarter - a self-hosted platform for managing projects, git repositories, and development tools.
|
||||
|
||||
## Overview
|
||||
|
||||
Built with **FastAPI** and **SQLAlchemy** (async), using **PostgreSQL** for data storage and **Docker** for tool instance management.
|
||||
|
||||
### Tech Stack
|
||||
|
||||
- **Framework**: FastAPI (Python 3.12+)
|
||||
- **Database**: PostgreSQL 15+ with asyncpg
|
||||
- **ORM**: SQLAlchemy 2.0 (async)
|
||||
- **Auth**: OAuth2 via Authentik with session cookies
|
||||
- **Migrations**: Alembic
|
||||
- **Tools**: Docker Compose for instance management
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Prerequisites
|
||||
|
||||
- Python 3.12+
|
||||
- PostgreSQL 15+ running locally
|
||||
- Docker (for tool instances)
|
||||
|
||||
### Setup
|
||||
|
||||
```bash
|
||||
cd apps/api
|
||||
|
||||
# Create virtual environment
|
||||
python -m venv .venv
|
||||
source .venv/bin/activate
|
||||
|
||||
# Install dependencies
|
||||
pip install -e ".[dev]"
|
||||
|
||||
# Set up database
|
||||
# Ensure PostgreSQL is running with a 'headquarter' database
|
||||
|
||||
# Run migrations
|
||||
alembic upgrade head
|
||||
|
||||
# Start development server
|
||||
uvicorn src.main:app --reload --port 8000
|
||||
```
|
||||
|
||||
The API will be available at `http://localhost:8000`.
|
||||
|
||||
### Interactive Documentation
|
||||
|
||||
Once running, visit:
|
||||
- **Swagger UI**: http://localhost:8000/docs
|
||||
- **ReDoc**: http://localhost:8000/redoc
|
||||
- **OpenAPI JSON**: http://localhost:8000/openapi.json
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Required | Default | Description |
|
||||
|----------|----------|---------|-------------|
|
||||
| `DATABASE_URL` | Yes | - | PostgreSQL connection string |
|
||||
| `API_BASE_URL` | Yes | - | Public API URL (e.g., `https://api.example.com`) |
|
||||
| `AUTHENTIK_DOMAIN` | Yes | - | Authentik server domain |
|
||||
| `AUTHENTIK_CLIENT_ID` | Yes | - | OAuth2 client ID |
|
||||
| `AUTHENTIK_CLIENT_SECRET` | Yes | - | OAuth2 client secret |
|
||||
| `AUTHENTIK_APPLICATION_SLUG` | Yes | - | Authentik application slug |
|
||||
| `WEB_BASE_URL` | Yes | - | Public frontend URL |
|
||||
| `SESSION_SECRET` | Yes | - | Secret for session cookie signing |
|
||||
| `COOKIE_DOMAIN` | No | - | Cookie domain (e.g., `.example.com`) |
|
||||
| `UPLOAD_DIR` | No | `./uploads` | Directory for file uploads |
|
||||
| `REPO_BASE_PATH` | No | `./repositories` | Base path for git repositories |
|
||||
| `INSTANCES_BASE_PATH` | No | `./instances` | Base path for tool instances |
|
||||
| `LOG_LEVEL` | No | `INFO` | Logging level |
|
||||
|
||||
## Development
|
||||
|
||||
### Running Tests
|
||||
|
||||
```bash
|
||||
# Run all tests
|
||||
pytest
|
||||
|
||||
# Run specific test category
|
||||
pytest -m unit # Unit tests (no DB)
|
||||
pytest -m integration # Integration tests (requires DB)
|
||||
|
||||
# Run with coverage
|
||||
pytest --cov=src --cov-report=html
|
||||
```
|
||||
|
||||
### Code Quality
|
||||
|
||||
```bash
|
||||
# Format code
|
||||
ruff format src tests
|
||||
|
||||
# Lint
|
||||
ruff check src tests
|
||||
|
||||
# Type check
|
||||
mypy src
|
||||
```
|
||||
|
||||
### Database Migrations
|
||||
|
||||
```bash
|
||||
# Create new migration
|
||||
alembic revision --autogenerate -m "description"
|
||||
|
||||
# Apply migrations
|
||||
alembic upgrade head
|
||||
|
||||
# Rollback one migration
|
||||
alembic downgrade -1
|
||||
|
||||
# Show current revision
|
||||
alembic current
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
### Directory Structure
|
||||
|
||||
```
|
||||
src/
|
||||
├── api/ # API endpoint routers
|
||||
│ ├── auth.py # OAuth2 authentication
|
||||
│ ├── dashboard.py # Dashboard summary
|
||||
│ ├── git_repositories.py # Git repo management
|
||||
│ ├── health.py # Health checks
|
||||
│ ├── projects.py # Project CRUD
|
||||
│ ├── ssh_keys.py # SSH key management
|
||||
│ ├── terminal.py # WebSocket terminal
|
||||
│ ├── tool_instances.py # Tool instance management
|
||||
│ ├── tool_types.py # Tool type definitions
|
||||
│ ├── user_config.py # User preferences
|
||||
│ └── users.py # User profile
|
||||
├── auth/ # Authentication logic
|
||||
│ ├── cookies.py # Cookie utilities
|
||||
│ ├── dependencies.py # Auth dependencies
|
||||
│ ├── oidc.py # OpenID Connect
|
||||
│ └── session.py # Session management
|
||||
├── config.py # Application settings
|
||||
├── database.py # Database setup
|
||||
├── main.py # FastAPI application
|
||||
├── models/ # SQLAlchemy models
|
||||
├── schemas/ # Pydantic schemas
|
||||
├── services/ # Business logic
|
||||
│ ├── docker.py # Docker Compose management
|
||||
│ ├── terminal_manager.py # Terminal sessions
|
||||
│ └── terminal_session.py # Terminal I/O
|
||||
└── utils/ # Utilities
|
||||
├── git_control.py # Git operations
|
||||
├── git_files.py # File operations
|
||||
├── git_history.py # History extraction
|
||||
└── git_url_parser.py # URL parsing
|
||||
```
|
||||
|
||||
### Authentication Flow
|
||||
|
||||
1. User clicks "Login" → redirects to Authentik OAuth
|
||||
2. Authentik redirects back with authorization code
|
||||
3. API exchanges code for tokens and fetches user info
|
||||
4. API creates session cookie (HMAC-signed, httpOnly)
|
||||
5. Frontend stores nothing - cookie sent automatically
|
||||
6. Subsequent requests include cookie for authentication
|
||||
|
||||
### Data Flow
|
||||
|
||||
```
|
||||
Client → FastAPI Router → Auth Dependency → Service Layer → Database
|
||||
↓
|
||||
Pydantic Models (validation)
|
||||
↓
|
||||
SQLAlchemy Models (ORM)
|
||||
↓
|
||||
PostgreSQL (storage)
|
||||
```
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Authentication
|
||||
- `GET /auth/login` - Initiate OAuth login
|
||||
- `GET /auth/callback` - OAuth callback
|
||||
- `GET /auth/me` - Get current user
|
||||
- `POST /auth/logout` - Logout
|
||||
|
||||
### Projects
|
||||
- `GET /projects` - List projects
|
||||
- `POST /projects` - Create project
|
||||
- `GET /projects/{id}` - Get project
|
||||
- `PUT /projects/{id}` - Update project
|
||||
- `DELETE /projects/{id}` - Delete project
|
||||
|
||||
### Git Repositories
|
||||
- `GET /projects/{id}/repositories` - List repositories
|
||||
- `POST /projects/{id}/repositories` - Create repository
|
||||
- `GET /projects/{id}/repositories/{id}` - Get repository
|
||||
- `DELETE /projects/{id}/repositories/{id}` - Delete repository
|
||||
- `GET /projects/{id}/repositories/{id}/files` - List files
|
||||
- `GET /projects/{id}/repositories/{id}/files/content` - Get file content
|
||||
- `POST /projects/{id}/repositories/{id}/files/content` - Update file
|
||||
- `GET /projects/{id}/repositories/{id}/branches` - List branches
|
||||
- `GET /projects/{id}/repositories/{id}/history` - Commit history
|
||||
- `GET /projects/{id}/repositories/{id}/commits/{hash}` - Commit detail
|
||||
|
||||
### Tool Types
|
||||
- `GET /tool-types` - List tool types
|
||||
- `POST /tool-types` - Create tool type
|
||||
- `GET /tool-types/{id}` - Get tool type
|
||||
- `PUT /tool-types/{id}` - Update tool type
|
||||
- `DELETE /tool-types/{id}` - Delete tool type
|
||||
|
||||
### Tool Instances
|
||||
- `GET /tool-instances` - List instances
|
||||
- `POST /tool-instances` - Create instance
|
||||
- `GET /tool-instances/{id}` - Get instance
|
||||
- `POST /tool-instances/{id}/start` - Start instance
|
||||
- `POST /tool-instances/{id}/stop` - Stop instance
|
||||
- `POST /tool-instances/{id}/restart` - Restart instance
|
||||
- `DELETE /tool-instances/{id}` - Delete instance
|
||||
- `GET /tool-instances/{id}/logs` - Get logs
|
||||
|
||||
### Terminal
|
||||
- `WS /ws/tool-instances/{id}/terminal` - WebSocket terminal
|
||||
|
||||
### Users
|
||||
- `GET /users/me` - Get profile
|
||||
- `PUT /users/me` - Update profile
|
||||
- `POST /users/me/avatar` - Upload avatar
|
||||
- `GET /users/me/config` - Get config
|
||||
- `PATCH /users/me/config` - Update config
|
||||
|
||||
### SSH Keys
|
||||
- `GET /ssh-keys` - List keys
|
||||
- `POST /ssh-keys` - Create key
|
||||
- `DELETE /ssh-keys/{id}` - Delete key
|
||||
|
||||
### Health
|
||||
- `GET /health` - System health
|
||||
- `GET /health/db` - Database health
|
||||
|
||||
## Deployment
|
||||
|
||||
See the [deployment documentation](../../docs/deployment/) for Docker and Traefik setup.
|
||||
|
||||
## Contributing
|
||||
|
||||
1. Follow PEP 8 style guide
|
||||
2. Add tests for new endpoints
|
||||
3. Update documentation
|
||||
4. Run quality gates before committing
|
||||
@@ -11,8 +11,8 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '0003'
|
||||
down_revision: Union[str, None] = '0002'
|
||||
revision: str = '0003_user_configs'
|
||||
down_revision: Union[str, None] = '0002_refresh_tokens'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
@@ -27,7 +27,8 @@ def upgrade() -> None:
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('user_id')
|
||||
sa.UniqueConstraint('user_id'),
|
||||
if_not_exists=True,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
"""add tool_types table
|
||||
|
||||
Revision ID: 0004_tool_types
|
||||
Revises: 0003_user_configs
|
||||
Create Date: 2026-05-18 15:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0004_tool_types"
|
||||
down_revision: Union[str, None] = "0003_user_configs"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"tool_types",
|
||||
sa.Column("id", sa.Uuid(as_uuid=True), primary_key=True),
|
||||
sa.Column("name", sa.String(255), nullable=False, unique=True),
|
||||
sa.Column("display_name", sa.String(255), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("compose_template", sa.Text(), nullable=False),
|
||||
sa.Column("required_variables", sa.JSON(), nullable=False, default=list),
|
||||
sa.Column("is_builtin", sa.Boolean(), nullable=False, default=False),
|
||||
sa.Column("created_by_id", sa.Uuid(as_uuid=True), sa.ForeignKey("users.id"), nullable=True),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("now()"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("now()"),
|
||||
onupdate=sa.text("now()"),
|
||||
nullable=False,
|
||||
),
|
||||
if_not_exists=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("tool_types")
|
||||
@@ -0,0 +1,44 @@
|
||||
"""add timestamps to ssh_keys table
|
||||
|
||||
Revision ID: 0005_ssh_keys_timestamps
|
||||
Revises: 0004_tool_types
|
||||
Create Date: 2026-05-19 09:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0005_ssh_keys_timestamps"
|
||||
down_revision: Union[str, None] = "0004_tool_types"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"ssh_keys",
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("now()"),
|
||||
nullable=True,
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"ssh_keys",
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("now()"),
|
||||
nullable=True,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("ssh_keys", "updated_at")
|
||||
op.drop_column("ssh_keys", "created_at")
|
||||
@@ -0,0 +1,55 @@
|
||||
"""add tool_instances table
|
||||
|
||||
Revision ID: 0006_tool_instances
|
||||
Revises: 0005_ssh_keys_timestamps
|
||||
Create Date: 2026-05-19 10:00:00.000000
|
||||
|
||||
"""
|
||||
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 = "0006_tool_instances"
|
||||
down_revision: Union[str, None] = "0005_ssh_keys_timestamps"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"tool_instances",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("name", sa.String(255), nullable=False),
|
||||
sa.Column("display_name", sa.String(255), nullable=False),
|
||||
sa.Column("tool_type_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("repository_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("project_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("owner_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("status", sa.String(50), nullable=False, server_default="pending"),
|
||||
sa.Column("container_id", sa.String(255), nullable=True),
|
||||
sa.Column("compose_path", sa.String(1024), nullable=True),
|
||||
sa.Column("url", sa.String(1024), nullable=True),
|
||||
sa.Column("port", sa.Integer(), nullable=True),
|
||||
sa.Column("last_started_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("last_stopped_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.ForeignKeyConstraint(["tool_type_id"], ["tool_types.id"]),
|
||||
sa.ForeignKeyConstraint(["repository_id"], ["git_repositories.id"]),
|
||||
sa.ForeignKeyConstraint(["project_id"], ["projects.id"]),
|
||||
sa.ForeignKeyConstraint(["owner_id"], ["users.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("idx_tool_instances_owner", "tool_instances", ["owner_id"])
|
||||
op.create_index("idx_tool_instances_repo", "tool_instances", ["repository_id"])
|
||||
op.create_index("idx_tool_instances_status", "tool_instances", ["status"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("idx_tool_instances_status", table_name="tool_instances")
|
||||
op.drop_index("idx_tool_instances_repo", table_name="tool_instances")
|
||||
op.drop_index("idx_tool_instances_owner", table_name="tool_instances")
|
||||
op.drop_table("tool_instances")
|
||||
@@ -0,0 +1,28 @@
|
||||
"""add container_name to tool_instances
|
||||
|
||||
Revision ID: 0007_instance_container_name
|
||||
Revises: 0006_tool_instances
|
||||
Create Date: 2026-05-20 08:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0007_instance_container_name"
|
||||
down_revision: Union[str, None] = "0006_tool_instances"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_instances",
|
||||
sa.Column("container_name", sa.String(255), nullable=True)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_instances", "container_name")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""add category and interfaces to tool_types
|
||||
|
||||
Revision ID: 0008_tool_type_category
|
||||
Revises: 0007_instance_container_name
|
||||
Create Date: 2026-05-20 09:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0008_tool_type_category"
|
||||
down_revision: Union[str, None] = "0007_instance_container_name"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_types",
|
||||
sa.Column("category", sa.String(50), nullable=False, server_default="other")
|
||||
)
|
||||
op.add_column(
|
||||
"tool_types",
|
||||
sa.Column("interfaces", sa.JSON(), nullable=False, server_default='["web"]')
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_types", "interfaces")
|
||||
op.drop_column("tool_types", "category")
|
||||
@@ -0,0 +1,46 @@
|
||||
"""add tool_configs table
|
||||
|
||||
Revision ID: 0009_tool_configs
|
||||
Revises: 0008_tool_type_category
|
||||
Create Date: 2026-05-20 09:00:00.000000
|
||||
|
||||
"""
|
||||
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 = "0009_tool_configs"
|
||||
down_revision: Union[str, None] = "0008_tool_type_category"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"tool_configs",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("tool_type_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("project_id", postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column("key", sa.String(255), nullable=False),
|
||||
sa.Column("value", sa.Text(), nullable=False),
|
||||
sa.Column("config_type", sa.String(20), nullable=False, server_default="env"),
|
||||
sa.Column("file_path", sa.String(1024), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
|
||||
sa.ForeignKeyConstraint(["tool_type_id"], ["tool_types.id"]),
|
||||
sa.ForeignKeyConstraint(["project_id"], ["projects.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("idx_tool_configs_user_tool", "tool_configs", ["user_id", "tool_type_id"])
|
||||
op.create_index("idx_tool_configs_project", "tool_configs", ["project_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("idx_tool_configs_project", table_name="tool_configs")
|
||||
op.drop_index("idx_tool_configs_user_tool", table_name="tool_configs")
|
||||
op.drop_table("tool_configs")
|
||||
@@ -0,0 +1,28 @@
|
||||
"""add default_port to tool_types
|
||||
|
||||
Revision ID: 0010_tool_type_default_port
|
||||
Revises: 0009_tool_configs
|
||||
Create Date: 2026-05-20 10:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0010_tool_type_default_port"
|
||||
down_revision: Union[str, None] = "0009_tool_configs"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_types",
|
||||
sa.Column("default_port", sa.Integer(), nullable=True)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_types", "default_port")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""add tunnel fields to tool_instances
|
||||
|
||||
Revision ID: 0011_tool_instance_tunnel_fields
|
||||
Revises: 0010_tool_type_default_port
|
||||
Create Date: 2026-05-20 12:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0011_tool_instance_tunnel_fields"
|
||||
down_revision: Union[str, None] = "0010_tool_type_default_port"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_instances",
|
||||
sa.Column("public_url", sa.String(1024), nullable=True)
|
||||
)
|
||||
op.add_column(
|
||||
"tool_instances",
|
||||
sa.Column("tunnel_id", sa.String(255), nullable=True)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_instances", "tunnel_id")
|
||||
op.drop_column("tool_instances", "public_url")
|
||||
@@ -0,0 +1,48 @@
|
||||
"""make default_port non-nullable and set values
|
||||
|
||||
Revision ID: 0012_default_port_req
|
||||
Revises: 0011_tool_instance_tunnel_fields
|
||||
Create Date: 2026-05-20 15:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0012_default_port_req"
|
||||
down_revision: Union[str, None] = "0011_tool_instance_tunnel_fields"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Set default_port for existing built-in tool types
|
||||
op.execute("""
|
||||
UPDATE tool_types
|
||||
SET default_port = CASE
|
||||
WHEN name = 'code-server' THEN 8443
|
||||
WHEN name = 'jupyter-notebook' THEN 8888
|
||||
WHEN name = 'opencode' THEN 3000
|
||||
ELSE 8080
|
||||
END
|
||||
WHERE default_port IS NULL
|
||||
""")
|
||||
|
||||
# Make default_port non-nullable
|
||||
op.alter_column(
|
||||
"tool_types",
|
||||
"default_port",
|
||||
existing_type=sa.Integer(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.alter_column(
|
||||
"tool_types",
|
||||
"default_port",
|
||||
existing_type=sa.Integer(),
|
||||
nullable=True,
|
||||
)
|
||||
@@ -0,0 +1,29 @@
|
||||
"""add probe_result to tool_instances
|
||||
|
||||
Revision ID: 0013_add_probe_result
|
||||
Revises: 0012_default_port_req
|
||||
Create Date: 2026-05-22 21:45:00.000000
|
||||
|
||||
"""
|
||||
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 = "0013_add_probe_result"
|
||||
down_revision: Union[str, None] = "0012_default_port_req"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_instances",
|
||||
sa.Column("probe_result", postgresql.JSON, nullable=True)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_instances", "probe_result")
|
||||
@@ -0,0 +1,23 @@
|
||||
"""merge migration heads
|
||||
|
||||
Revision ID: 0014_merge_heads
|
||||
Revises: 0013_add_probe_result, 8ed7dd80973d
|
||||
Create Date: 2026-05-22 21:50:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0014_merge_heads"
|
||||
down_revision: Union[str, Sequence[str], None] = ("0013_add_probe_result", "8ed7dd80973d")
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -0,0 +1,109 @@
|
||||
"""replace interfaces with interface_type and add requires_port
|
||||
|
||||
Revision ID: 0015_single_interface
|
||||
Revises: 0014_merge_heads
|
||||
Create Date: 2026-05-22 22:00:00.000000
|
||||
|
||||
"""
|
||||
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 = "0015_single_interface"
|
||||
down_revision: Union[str, Sequence[str], None] = "0014_merge_heads"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _get_dialect() -> str:
|
||||
"""Get the current database dialect name."""
|
||||
conn = op.get_bind()
|
||||
return conn.dialect.name
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
dialect = _get_dialect()
|
||||
|
||||
# Add new columns
|
||||
op.add_column('tool_types', sa.Column('interface_type', sa.String(20), nullable=True))
|
||||
op.add_column('tool_types', sa.Column('requires_port', sa.Boolean(), nullable=False, server_default='true'))
|
||||
|
||||
# Migrate data: take first element from interfaces JSON array
|
||||
if dialect == 'postgresql':
|
||||
op.execute("""
|
||||
UPDATE tool_types
|
||||
SET interface_type = COALESCE(
|
||||
(SELECT elem FROM jsonb_array_elements_text(interfaces::jsonb) AS elem LIMIT 1),
|
||||
'web'
|
||||
),
|
||||
requires_port = CASE
|
||||
WHEN COALESCE(
|
||||
(SELECT elem FROM jsonb_array_elements_text(interfaces::jsonb) AS elem LIMIT 1),
|
||||
'web'
|
||||
) = 'web' THEN true
|
||||
ELSE false
|
||||
END
|
||||
""")
|
||||
else:
|
||||
# SQLite: interfaces is stored as JSON text, extract first array element
|
||||
op.execute("""
|
||||
UPDATE tool_types
|
||||
SET interface_type = COALESCE(
|
||||
(SELECT json_extract(value, '$[0]')
|
||||
FROM json_each(interfaces) AS value
|
||||
WHERE json_valid(interfaces)
|
||||
LIMIT 1),
|
||||
'web'
|
||||
),
|
||||
requires_port = CASE
|
||||
WHEN COALESCE(
|
||||
(SELECT json_extract(value, '$[0]')
|
||||
FROM json_each(interfaces) AS value
|
||||
WHERE json_valid(interfaces)
|
||||
LIMIT 1),
|
||||
'web'
|
||||
) = 'web' THEN true
|
||||
ELSE false
|
||||
END
|
||||
""")
|
||||
|
||||
# Make interface_type non-nullable after data migration
|
||||
op.alter_column('tool_types', 'interface_type', nullable=False)
|
||||
|
||||
# Drop old interfaces column
|
||||
op.drop_column('tool_types', 'interfaces')
|
||||
|
||||
# Add CHECK constraint for interface_type (only on PostgreSQL; SQLite supports it too)
|
||||
op.create_check_constraint('chk_interface_type', 'tool_types', sa.text("interface_type IN ('web', 'terminal')"))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
dialect = _get_dialect()
|
||||
|
||||
# Drop CHECK constraint
|
||||
op.drop_constraint('chk_interface_type', 'tool_types', type_='check')
|
||||
|
||||
# Add back interfaces column
|
||||
if dialect == 'postgresql':
|
||||
op.add_column('tool_types', sa.Column('interfaces', postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default='["web"]'))
|
||||
|
||||
# Migrate data back: wrap interface_type in array
|
||||
op.execute("""
|
||||
UPDATE tool_types
|
||||
SET interfaces = jsonb_build_array(interface_type)
|
||||
""")
|
||||
else:
|
||||
op.add_column('tool_types', sa.Column('interfaces', sa.JSON(), nullable=False, server_default='["web"]'))
|
||||
|
||||
# Migrate data back: wrap interface_type in array for SQLite
|
||||
op.execute("""
|
||||
UPDATE tool_types
|
||||
SET interfaces = json_array(interface_type)
|
||||
""")
|
||||
|
||||
# Drop new columns
|
||||
op.drop_column('tool_types', 'requires_port')
|
||||
op.drop_column('tool_types', 'interface_type')
|
||||
@@ -0,0 +1,129 @@
|
||||
"""add pi agent tool type
|
||||
|
||||
Revision ID: 20260527_160017_add_pi_agent
|
||||
Revises: f3d2dc90ba3a
|
||||
Create Date: 2026-05-27T16:00:17
|
||||
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
import uuid
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "20260527_160017_add_pi_agent"
|
||||
down_revision: Union[str, None] = "2026_05_27_external_repos"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
PI_AGENT_ID = uuid.UUID("d07b8376-2151-4119-8c1d-27f792aae9a3")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Check if pi-agent already exists
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(
|
||||
sa.text("SELECT id FROM tool_types WHERE name = 'pi-agent'")
|
||||
).fetchone()
|
||||
|
||||
if result is None:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO tool_types (
|
||||
id, name, display_name, description, category,
|
||||
interface_type, requires_port, default_port,
|
||||
definition_type, compose_template, dockerfile_template, required_variables,
|
||||
created_at, updated_at
|
||||
) VALUES (
|
||||
:id, :name, :display_name, :description, :category,
|
||||
:interface_type, :requires_port, :default_port,
|
||||
:definition_type, :compose_template, :dockerfile_template, :required_variables,
|
||||
now(), now()
|
||||
)
|
||||
"""),
|
||||
{
|
||||
"id": PI_AGENT_ID,
|
||||
"name": "pi-agent",
|
||||
"display_name": "Pi Agent",
|
||||
"description": "Pi coding agent terminal environment with nvim, ranger, and tmux",
|
||||
"category": "development",
|
||||
"interface_type": "terminal",
|
||||
"requires_port": False,
|
||||
"default_port": 0,
|
||||
"definition_type": "dockerfile",
|
||||
"compose_template": """services:
|
||||
app:
|
||||
build: .
|
||||
stdin_open: true
|
||||
tty: true
|
||||
volumes:
|
||||
- ${REPO_PATH}:/workspace
|
||||
working_dir: /workspace
|
||||
command: /bin/bash""",
|
||||
"dockerfile_template": """# Pi Coding Agent - Terminal-based coding harness
|
||||
FROM ubuntu:24.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install base dependencies
|
||||
RUN apt-get update && apt-get install -y \\
|
||||
curl \\
|
||||
wget \\
|
||||
git \\
|
||||
neovim \\
|
||||
ranger \\
|
||||
tmux \\
|
||||
htop \\
|
||||
tree \\
|
||||
jq \\
|
||||
ca-certificates \\
|
||||
python3 \\
|
||||
python3-pip \\
|
||||
build-essential \\
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install Node.js (required for Pi)
|
||||
RUN curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \\
|
||||
&& apt-get install -y nodejs \\
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install Pi Coding Agent globally
|
||||
RUN npm install -g --ignore-scripts @earendil-works/pi-coding-agent
|
||||
|
||||
# Create non-root user
|
||||
RUN useradd -m -s /bin/bash user
|
||||
WORKDIR /home/user
|
||||
|
||||
# Set up git
|
||||
RUN git config --global init.defaultBranch main \\
|
||||
&& git config --global user.email "dev@headquarter.local" \\
|
||||
&& git config --global user.name "Developer"
|
||||
|
||||
# Create default tmux config
|
||||
RUN echo 'set -g mouse on\\nset -g default-terminal "screen-256color"' > /home/user/.tmux.conf
|
||||
|
||||
# Create default ranger config
|
||||
RUN mkdir -p /home/user/.config/ranger \\
|
||||
&& echo 'set preview_files true\\nset use_preview_script true' > /home/user/.config/ranger/rc.conf
|
||||
|
||||
# Set up Pi config directory
|
||||
RUN mkdir -p /home/user/.pi/agent
|
||||
|
||||
USER user
|
||||
|
||||
# Default to bash (Pi is invoked manually via `pi` command)
|
||||
CMD ["/bin/bash"]""",
|
||||
"required_variables": json.dumps(["REPO_PATH"]),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
conn.execute(
|
||||
sa.text("DELETE FROM tool_types WHERE name = 'pi-agent'")
|
||||
)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""add_clone_mode_and_ssh_key_id
|
||||
|
||||
Revision ID: 2026_05_22_add_clone_mode
|
||||
Revises: 0014_merge_heads
|
||||
Create Date: 2026-05-22 20:30:00.000000
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '2026_05_22_add_clone_mode'
|
||||
down_revision = '0015_single_interface'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add ssh_key_id to git_repositories
|
||||
op.add_column('git_repositories', sa.Column('ssh_key_id', postgresql.UUID(), nullable=True))
|
||||
op.create_foreign_key('fk_git_repositories_ssh_key', 'git_repositories', 'ssh_keys', ['ssh_key_id'], ['id'])
|
||||
|
||||
# Add clone_mode and branch to tool_instances
|
||||
op.add_column('tool_instances', sa.Column('clone_mode', sa.String(20), nullable=False, server_default='mount'))
|
||||
op.add_column('tool_instances', sa.Column('branch', sa.String(255), nullable=True, server_default='main'))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop columns from tool_instances
|
||||
op.drop_column('tool_instances', 'branch')
|
||||
op.drop_column('tool_instances', 'clone_mode')
|
||||
|
||||
# Drop ssh_key_id from git_repositories
|
||||
op.drop_constraint('fk_git_repositories_ssh_key', 'git_repositories', type_='foreignkey')
|
||||
op.drop_column('git_repositories', 'ssh_key_id')
|
||||
@@ -0,0 +1,25 @@
|
||||
"""remove_is_builtin_from_tool_types
|
||||
|
||||
Revision ID: 2026_05_23_remove_is_builtin
|
||||
Revises: 2026_05_22_add_clone_mode
|
||||
Create Date: 2026-05-23 14:30:00.000000
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '2026_05_23_remove_is_builtin'
|
||||
down_revision = 'f3d2dc90ba3a'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Drop the is_builtin column from tool_types
|
||||
op.execute("ALTER TABLE tool_types DROP COLUMN IF EXISTS is_builtin")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Add the is_builtin column back to tool_types
|
||||
op.add_column('tool_types', sa.Column('is_builtin', sa.Boolean(), nullable=False, server_default='false'))
|
||||
@@ -0,0 +1,30 @@
|
||||
"""add startup_command to tool_types
|
||||
|
||||
Revision ID: 2026_05_24_220141
|
||||
Revises: 6fc7bfcf199f
|
||||
Create Date: 2026-05-24 22:01:41.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "2026_05_24_220141"
|
||||
down_revision: Union[str, Sequence[str], None] = "6fc7bfcf199f"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"tool_types",
|
||||
sa.Column("startup_command", sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("tool_types", "startup_command")
|
||||
@@ -0,0 +1,86 @@
|
||||
"""add_config_profiles
|
||||
|
||||
Revision ID: 2026_05_24_add_config_profiles
|
||||
Revises: f3d2dc90ba3a
|
||||
Create Date: 2026-05-24 14:00:00.000000
|
||||
"""
|
||||
|
||||
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 = "2026_05_24_add_config_profiles"
|
||||
down_revision: Union[str, Sequence[str], None] = "f3d2dc90ba3a"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Create config_profiles table
|
||||
op.create_table(
|
||||
"config_profiles",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False),
|
||||
sa.Column("name", sa.String(255), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("project_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("projects.id", ondelete="CASCADE"), nullable=True),
|
||||
sa.Column("tool_type_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("tool_types.id", ondelete="CASCADE"), nullable=True),
|
||||
sa.Column("env_vars", postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default="{}"),
|
||||
sa.Column("runtime_hints", postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default="{}"),
|
||||
sa.Column("mounts", postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default="[]"),
|
||||
sa.Column("files", postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default="{}"),
|
||||
sa.Column("is_default", sa.Boolean(), nullable=False, server_default="false"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.text("NOW()"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.text("NOW()"), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("user_id", "name", name="uq_config_profiles_user_name"),
|
||||
)
|
||||
|
||||
# Create indexes for config_profiles
|
||||
op.create_index("idx_config_profiles_user", "config_profiles", ["user_id"])
|
||||
op.create_index("idx_config_profiles_project", "config_profiles", ["project_id"])
|
||||
op.create_index("idx_config_profiles_tool_type", "config_profiles", ["tool_type_id"])
|
||||
|
||||
# Create config_profile_includes table
|
||||
op.create_table(
|
||||
"config_profile_includes",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("profile_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("config_profiles.id", ondelete="CASCADE"), nullable=False),
|
||||
sa.Column("included_profile_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("config_profiles.id", ondelete="CASCADE"), nullable=False),
|
||||
sa.Column("order_index", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.text("NOW()"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.text("NOW()"), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("profile_id", "included_profile_id", name="uq_config_profile_includes"),
|
||||
)
|
||||
|
||||
# Create indexes for config_profile_includes
|
||||
op.create_index("idx_config_profile_includes_profile", "config_profile_includes", ["profile_id"])
|
||||
op.create_index("idx_config_profile_includes_included", "config_profile_includes", ["included_profile_id"])
|
||||
|
||||
# Add selected_config_profile_id to tool_instances
|
||||
op.add_column(
|
||||
"tool_instances",
|
||||
sa.Column("selected_config_profile_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("config_profiles.id", ondelete="SET NULL"), nullable=True),
|
||||
)
|
||||
op.create_index("idx_tool_instances_config_profile", "tool_instances", ["selected_config_profile_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Remove selected_config_profile_id from tool_instances
|
||||
op.drop_index("idx_tool_instances_config_profile", table_name="tool_instances")
|
||||
op.drop_column("tool_instances", "selected_config_profile_id")
|
||||
|
||||
# Drop config_profile_includes table
|
||||
op.drop_index("idx_config_profile_includes_included", table_name="config_profile_includes")
|
||||
op.drop_index("idx_config_profile_includes_profile", table_name="config_profile_includes")
|
||||
op.drop_table("config_profile_includes")
|
||||
|
||||
# Drop config_profiles table
|
||||
op.drop_index("idx_config_profiles_tool_type", table_name="config_profiles")
|
||||
op.drop_index("idx_config_profiles_project", table_name="config_profiles")
|
||||
op.drop_index("idx_config_profiles_user", table_name="config_profiles")
|
||||
op.drop_table("config_profiles")
|
||||
@@ -0,0 +1,28 @@
|
||||
"""add_git_mounts_to_config_profiles
|
||||
|
||||
Revision ID: 2026_05_26_add_git_mounts
|
||||
Revises: f3d2dc90ba3a
|
||||
Create Date: 2026-05-26 12:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "2026_05_26_add_git_mounts"
|
||||
down_revision: Union[str, Sequence[str], None] = "2026_05_24_220141"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"config_profiles",
|
||||
sa.Column("git_mounts", sa.JSON(), nullable=True, default=list),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("config_profiles", "git_mounts")
|
||||
@@ -0,0 +1,41 @@
|
||||
"""make_project_id_nullable_in_git_repositories
|
||||
|
||||
Revision ID: 2026_05_27_external_repos
|
||||
Revises: 2026_05_26_add_git_mounts
|
||||
Create Date: 2026-05-27 08:30:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "2026_05_27_external_repos"
|
||||
down_revision: Union[str, Sequence[str], None] = "2026_05_26_add_git_mounts"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Expand alembic_version version_num to avoid truncation errors
|
||||
op.execute("ALTER TABLE alembic_version ALTER COLUMN version_num TYPE VARCHAR(64)")
|
||||
|
||||
# Make project_id nullable to allow external repositories
|
||||
op.alter_column(
|
||||
"git_repositories",
|
||||
"project_id",
|
||||
existing_type=sa.UUID(),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.alter_column(
|
||||
"git_repositories",
|
||||
"project_id",
|
||||
existing_type=sa.UUID(),
|
||||
nullable=False,
|
||||
)
|
||||
op.execute("ALTER TABLE alembic_version ALTER COLUMN version_num TYPE VARCHAR(32)")
|
||||
@@ -0,0 +1,40 @@
|
||||
"""add_tool_config_fields
|
||||
|
||||
Revision ID: 398082499c30
|
||||
Revises: af8512103d67
|
||||
Create Date: 2026-05-22 18:38:20.166184
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '398082499c30'
|
||||
down_revision = 'af8512103d67'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add new columns to tool_configs
|
||||
op.add_column('tool_configs', sa.Column('port_override', sa.Integer(), nullable=True))
|
||||
op.add_column('tool_configs', sa.Column('start_command', sa.Text(), nullable=True))
|
||||
op.add_column('tool_configs', sa.Column('working_directory', sa.Text(), nullable=True))
|
||||
op.add_column('tool_configs', sa.Column('environment_variables', postgresql.JSONB(astext_type=sa.Text()), nullable=True, server_default='{}'))
|
||||
op.add_column('tool_configs', sa.Column('volumes', postgresql.JSONB(astext_type=sa.Text()), nullable=True, server_default='[]'))
|
||||
|
||||
# Add CHECK constraint for port range
|
||||
op.create_check_constraint('chk_port_range', 'tool_configs', sa.text('port_override IS NULL OR (port_override >= 1 AND port_override <= 65535)'))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop CHECK constraint
|
||||
op.drop_constraint('chk_port_range', 'tool_configs', type_='check')
|
||||
|
||||
# Drop columns
|
||||
op.drop_column('tool_configs', 'port_override')
|
||||
op.drop_column('tool_configs', 'start_command')
|
||||
op.drop_column('tool_configs', 'working_directory')
|
||||
op.drop_column('tool_configs', 'environment_variables')
|
||||
op.drop_column('tool_configs', 'volumes')
|
||||
@@ -0,0 +1,23 @@
|
||||
"""merge_remove_is_builtin_and_add_config_profiles
|
||||
|
||||
Revision ID: 6fc7bfcf199f
|
||||
Revises: 2026_05_23_remove_is_builtin, 2026_05_24_add_config_profiles
|
||||
Create Date: 2026-05-24 18:00:43.990361
|
||||
"""
|
||||
|
||||
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '6fc7bfcf199f'
|
||||
down_revision = ('2026_05_23_remove_is_builtin', '2026_05_24_add_config_profiles')
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -0,0 +1,44 @@
|
||||
"""create_config_folders_table
|
||||
|
||||
Revision ID: 8ed7dd80973d
|
||||
Revises: 398082499c30
|
||||
Create Date: 2026-05-22 18:38:22.133696
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '8ed7dd80973d'
|
||||
down_revision = '398082499c30'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'config_folders',
|
||||
sa.Column('id', postgresql.UUID(as_uuid=True), primary_key=True, server_default=sa.text('gen_random_uuid()')),
|
||||
sa.Column('user_id', postgresql.UUID(as_uuid=True), sa.ForeignKey('users.id', ondelete='CASCADE'), nullable=False),
|
||||
sa.Column('name', sa.String(255), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('mount_path', sa.String(1024), nullable=False),
|
||||
sa.Column('files', postgresql.JSONB(astext_type=sa.Text()), nullable=False, server_default='{}'),
|
||||
sa.Column('project_overrides', postgresql.JSONB(astext_type=sa.Text()), nullable=True, server_default='{}'),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=False, server_default='true'),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('NOW()')),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('NOW()')),
|
||||
sa.UniqueConstraint('user_id', 'name', name='uq_config_folders_user_name')
|
||||
)
|
||||
|
||||
# Add index on user_id for filtering
|
||||
op.create_index('idx_config_folders_user', 'config_folders', ['user_id'])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop index
|
||||
op.drop_index('idx_config_folders_user', table_name='config_folders')
|
||||
|
||||
# Drop table
|
||||
op.drop_table('config_folders')
|
||||
@@ -0,0 +1,38 @@
|
||||
"""add_tool_type_fields
|
||||
|
||||
Revision ID: af8512103d67
|
||||
Revises: 0012_default_port_req
|
||||
Create Date: 2026-05-22 18:37:56.607240
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'af8512103d67'
|
||||
down_revision = '0012_default_port_req'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add new columns to tool_types
|
||||
op.add_column('tool_types', sa.Column('definition_type', sa.String(20), nullable=False, server_default='compose'))
|
||||
op.add_column('tool_types', sa.Column('dockerfile_template', sa.Text(), nullable=True))
|
||||
op.add_column('tool_types', sa.Column('build_context', postgresql.JSONB(astext_type=sa.Text()), nullable=True, server_default='{}'))
|
||||
op.add_column('tool_types', sa.Column('readiness_probe', postgresql.JSONB(astext_type=sa.Text()), nullable=True))
|
||||
|
||||
# Add CHECK constraint for definition_type
|
||||
op.create_check_constraint('chk_definition_type', 'tool_types', sa.text("definition_type IN ('compose', 'dockerfile')"))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop CHECK constraint
|
||||
op.drop_constraint('chk_definition_type', 'tool_types', type_='check')
|
||||
|
||||
# Drop columns
|
||||
op.drop_column('tool_types', 'definition_type')
|
||||
op.drop_column('tool_types', 'dockerfile_template')
|
||||
op.drop_column('tool_types', 'build_context')
|
||||
op.drop_column('tool_types', 'readiness_probe')
|
||||
@@ -0,0 +1,23 @@
|
||||
"""merge_single_interface_and_clone_mode
|
||||
|
||||
Revision ID: f3d2dc90ba3a
|
||||
Revises: 0015_single_interface, 2026_05_22_add_clone_mode
|
||||
Create Date: 2026-05-24 10:43:14.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f3d2dc90ba3a"
|
||||
down_revision: Union[str, Sequence[str], None] = ("0015_single_interface", "2026_05_22_add_clone_mode")
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -11,7 +11,6 @@ dependencies = [
|
||||
"alembic>=1.12.0",
|
||||
"pydantic>=2.5.0",
|
||||
"pydantic-settings>=2.1.0",
|
||||
"python-jose[cryptography]>=3.3.0",
|
||||
"python-multipart>=0.0.6",
|
||||
"httpx>=0.25.0",
|
||||
"structlog>=23.2.0",
|
||||
@@ -28,6 +27,9 @@ dev = [
|
||||
"aiosqlite>=0.19.0",
|
||||
]
|
||||
|
||||
[tool.mypy]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = ["."]
|
||||
asyncio_mode = "auto"
|
||||
|
||||
+131
-122
@@ -1,6 +1,6 @@
|
||||
import logging
|
||||
from secrets import token_urlsafe
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import AsyncGenerator, Literal, cast
|
||||
from typing import Any, AsyncGenerator, Literal, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, Response, status
|
||||
@@ -9,18 +9,14 @@ from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import (
|
||||
build_login_redirect_url,
|
||||
exchange_code_for_tokens,
|
||||
fetch_jwks,
|
||||
verify_provider_access_token,
|
||||
)
|
||||
from src.auth.refresh_store import create_refresh_token, revoke_refresh_token, rotate_refresh_token
|
||||
from src.auth.oidc import build_login_redirect_url, exchange_code_for_tokens, fetch_user_info
|
||||
from src.auth.session import create_session_cookie, decode_session_cookie
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.user import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
||||
@@ -29,8 +25,21 @@ async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
yield session
|
||||
|
||||
|
||||
@router.get("/login")
|
||||
async def login() -> RedirectResponse:
|
||||
@router.get(
|
||||
"/login",
|
||||
summary="Initiate OAuth login",
|
||||
description="Redirects to the configured OAuth provider (Authentik) to start the authentication flow.",
|
||||
response_class=RedirectResponse,
|
||||
)
|
||||
async def login(next: str = "/") -> RedirectResponse:
|
||||
"""Initiate OAuth2 login flow.
|
||||
|
||||
Args:
|
||||
next: URL to redirect to after successful authentication.
|
||||
|
||||
Returns:
|
||||
RedirectResponse to the OAuth provider's authorization endpoint.
|
||||
"""
|
||||
settings = Settings()
|
||||
redirect_uri = f"{settings.api_base_url}/auth/callback"
|
||||
state = token_urlsafe(24)
|
||||
@@ -38,10 +47,11 @@ async def login() -> RedirectResponse:
|
||||
settings=settings,
|
||||
redirect_uri=redirect_uri,
|
||||
state=state,
|
||||
nonce=token_urlsafe(16),
|
||||
)
|
||||
logger.debug("Auth login initiated: redirect_uri=%s, next=%s", redirect_uri, next)
|
||||
response = RedirectResponse(location)
|
||||
response.set_cookie("auth_state", state, httponly=True, samesite="lax")
|
||||
response.set_cookie("auth_next", next, httponly=True, samesite="lax")
|
||||
return response
|
||||
|
||||
|
||||
@@ -49,142 +59,141 @@ async def login() -> RedirectResponse:
|
||||
async def callback(
|
||||
code: str,
|
||||
state: str,
|
||||
response: Response,
|
||||
auth_state: str | None = Cookie(default=None),
|
||||
auth_next: str | None = Cookie(default="/"),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
) -> RedirectResponse:
|
||||
logger.debug("Auth callback received: code=%s... state=%s", code[:10] if code else "None", state[:10] if state else "None")
|
||||
|
||||
if auth_state is None or auth_state != state:
|
||||
logger.warning("State mismatch: cookie=%s, param=%s", auth_state, state)
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid state")
|
||||
|
||||
settings = Settings()
|
||||
redirect_uri = f"{settings.api_base_url}/auth/callback"
|
||||
logger.debug("Exchanging code for tokens (redirect_uri=%s)", redirect_uri)
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
token_payload = await exchange_code_for_tokens(
|
||||
settings=settings,
|
||||
code=code,
|
||||
redirect_uri=redirect_uri,
|
||||
client=client,
|
||||
)
|
||||
jwks = await fetch_jwks(settings=settings, client=client)
|
||||
try:
|
||||
token_payload = await exchange_code_for_tokens(
|
||||
settings=settings,
|
||||
code=code,
|
||||
redirect_uri=redirect_uri,
|
||||
client=client,
|
||||
)
|
||||
logger.info("Token exchange successful")
|
||||
except Exception as exc:
|
||||
logger.error("Token exchange failed: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=f"token exchange failed: {exc}")
|
||||
|
||||
try:
|
||||
user_info = await fetch_user_info(
|
||||
settings=settings,
|
||||
access_token=token_payload["access_token"],
|
||||
client=client,
|
||||
)
|
||||
logger.debug("User info fetched successfully")
|
||||
except Exception as exc:
|
||||
logger.error("User info fetch failed: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="failed to fetch user info")
|
||||
|
||||
provider_claims = verify_provider_access_token(
|
||||
settings=settings,
|
||||
token=token_payload["access_token"],
|
||||
jwks=jwks,
|
||||
)
|
||||
authentik_id = str(user_info.get("sub", ""))
|
||||
email = str(user_info.get("email", f"{authentik_id}@authentik.local"))
|
||||
name = str(user_info.get("name", email))
|
||||
logger.debug("User info: authentik_id=%s, email=%s, name=%s", authentik_id, email, name)
|
||||
|
||||
authentik_id = str(provider_claims["sub"])
|
||||
email = str(provider_claims.get("email", f"{authentik_id}@authentik.local"))
|
||||
name = str(provider_claims.get("name", email))
|
||||
|
||||
user = await session.scalar(select(User).where(User.authentik_id == authentik_id))
|
||||
if user is None:
|
||||
user = User(email=email, name=name, authentik_id=authentik_id, avatar_url=None)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
else:
|
||||
user.email = email
|
||||
user.name = name
|
||||
await session.commit()
|
||||
|
||||
access_token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(user.id),
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=settings.access_token_ttl_minutes),
|
||||
)
|
||||
refresh_token, _ = await create_refresh_token(
|
||||
session=session,
|
||||
user_id=user.id,
|
||||
expires_at=datetime.now(UTC) + timedelta(days=settings.refresh_token_ttl_days),
|
||||
user_agent=None,
|
||||
ip_address=None,
|
||||
)
|
||||
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
response.set_cookie("access_token", access_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.set_cookie("refresh_token", refresh_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.delete_cookie("auth_state", samesite="lax")
|
||||
|
||||
return {"sub": str(user.id), "email": user.email, "name": user.name}
|
||||
|
||||
|
||||
@router.post("/refresh")
|
||||
async def refresh(
|
||||
response: Response,
|
||||
refresh_token: str | None = Cookie(default=None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
if not refresh_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing refresh token")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
rotated_raw_token, rotated_record = await rotate_refresh_token(
|
||||
session=session,
|
||||
raw_token=refresh_token,
|
||||
user_agent=None,
|
||||
ip_address=None,
|
||||
)
|
||||
except ValueError as error:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=str(error)) from error
|
||||
|
||||
user = await session.get(User, rotated_record.user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid refresh token")
|
||||
|
||||
access_token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(user.id),
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=settings.access_token_ttl_minutes),
|
||||
)
|
||||
user = await session.scalar(select(User).where(User.authentik_id == authentik_id))
|
||||
if user is None:
|
||||
logger.debug("Creating new user: authentik_id=%s", authentik_id)
|
||||
user = User(email=email, name=name, authentik_id=authentik_id, avatar_url=None)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
logger.info("New user created: id=%s", user.id)
|
||||
else:
|
||||
logger.debug("Existing user found: id=%s, updating info", user.id)
|
||||
user.email = email
|
||||
user.name = name
|
||||
await session.commit()
|
||||
except Exception as exc:
|
||||
logger.error("Database error during user lookup/creation: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="database error")
|
||||
|
||||
# Create session cookie
|
||||
session_cookie = create_session_cookie(settings=settings, user_id=str(user.id))
|
||||
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
|
||||
response.set_cookie("access_token", access_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.set_cookie("refresh_token", rotated_raw_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
|
||||
return {"sub": str(user.id), "email": user.email, "name": user.name}
|
||||
cookie_domain = str(cookie_options["domain"]) if cookie_options.get("domain") else None
|
||||
|
||||
logger.info("Auth callback complete for user id=%s, redirecting to %s", user.id, auth_next)
|
||||
|
||||
# Redirect to frontend with the original next path
|
||||
redirect_url = f"{settings.web_base_url}{auth_next}"
|
||||
redirect_response = RedirectResponse(url=redirect_url)
|
||||
|
||||
redirect_response.set_cookie(
|
||||
"session",
|
||||
session_cookie,
|
||||
httponly=True,
|
||||
samesite=cookie_samesite,
|
||||
secure=cookie_secure,
|
||||
domain=cookie_domain,
|
||||
)
|
||||
redirect_response.delete_cookie("auth_state", samesite="lax", domain=cookie_domain)
|
||||
redirect_response.delete_cookie("auth_next", samesite="lax", domain=cookie_domain)
|
||||
|
||||
return redirect_response
|
||||
|
||||
|
||||
@router.post("/logout")
|
||||
async def logout(
|
||||
response: Response,
|
||||
refresh_token: str | None = Cookie(default=None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
async def logout(response: Response) -> dict[str, str]:
|
||||
settings = Settings()
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
cookie_domain = str(cookie_options["domain"]) if cookie_options.get("domain") else None
|
||||
|
||||
if refresh_token:
|
||||
try:
|
||||
await revoke_refresh_token(session=session, raw_token=refresh_token)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
response.delete_cookie("access_token", samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.delete_cookie("refresh_token", samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.delete_cookie("session", samesite=cookie_samesite, secure=cookie_secure, domain=cookie_domain)
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@router.get("/me")
|
||||
async def me(access_token: str | None = Cookie(default=None)) -> dict[str, str]:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
async def me(
|
||||
session_cookie: str | None = Cookie(default=None, alias="session"),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, Any]:
|
||||
logger.debug("Auth /me called, cookie present: %s", bool(session_cookie))
|
||||
|
||||
if not session_cookie:
|
||||
logger.warning("Auth /me: missing session cookie")
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
settings = Settings()
|
||||
logger.debug("Auth /me: cookie_domain=%s, cookie_secure=%s, cookie_samesite=%s",
|
||||
settings.cookie_domain, settings.cookie_secure, settings.cookie_samesite)
|
||||
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
user_id = payload["user_id"]
|
||||
logger.debug("Auth /me: decoded session for user_id=%s", user_id)
|
||||
except ValueError as exc:
|
||||
logger.warning("Auth /me: invalid session: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=str(exc))
|
||||
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
logger.warning("Auth /me: user not found for id=%s", user_id)
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
|
||||
logger.info("Auth /me: success for user=%s", user.email)
|
||||
return {
|
||||
"sub": str(claims["sub"]),
|
||||
"email": str(claims["email"]),
|
||||
"name": str(claims["name"]),
|
||||
"user": {
|
||||
"id": str(user.id),
|
||||
"email": user.email,
|
||||
"name": user.name,
|
||||
"avatar_url": user.avatar_url or "",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,337 @@
|
||||
"""Config folder API endpoints."""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.api.shared_validators import validate_files as _validate_files, validate_mount_path as _validate_mount_path
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.config_folder import ConfigFolder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/config-folders", tags=["config-folders"])
|
||||
|
||||
|
||||
class ConfigFolderCreate(BaseModel):
|
||||
name: str = Field(description="Folder name (unique per user)")
|
||||
description: str | None = Field(default=None, description="Optional description")
|
||||
mount_path: str = Field(description="Default mount path in container")
|
||||
files: dict = Field(default_factory=dict, description="Files as {path: content}")
|
||||
|
||||
@field_validator("mount_path")
|
||||
@classmethod
|
||||
def validate_mount_path(cls, v: str) -> str:
|
||||
return _validate_mount_path(v)
|
||||
|
||||
@field_validator("files")
|
||||
@classmethod
|
||||
def validate_files(cls, v: dict) -> dict:
|
||||
return _validate_files(v)
|
||||
|
||||
|
||||
class ConfigFolderUpdate(BaseModel):
|
||||
name: str | None = Field(default=None, description="Folder name")
|
||||
description: str | None = Field(default=None, description="Optional description")
|
||||
mount_path: str | None = Field(default=None, description="Default mount path")
|
||||
files: dict | None = Field(default=None, description="Files as {path: content}")
|
||||
is_active: bool | None = Field(default=None, description="Active/inactive toggle")
|
||||
|
||||
@field_validator("mount_path")
|
||||
@classmethod
|
||||
def validate_mount_path(cls, v: str | None) -> str | None:
|
||||
return _validate_mount_path(v)
|
||||
|
||||
@field_validator("files")
|
||||
@classmethod
|
||||
def validate_files(cls, v: dict | None) -> dict | None:
|
||||
return _validate_files(v)
|
||||
|
||||
|
||||
class ProjectOverrideCreate(BaseModel):
|
||||
mount_path: str | None = Field(default=None, description="Override mount path")
|
||||
files: dict = Field(default_factory=dict, description="Override files")
|
||||
|
||||
@field_validator("mount_path")
|
||||
@classmethod
|
||||
def validate_mount_path(cls, v: str | None) -> str | None:
|
||||
return _validate_mount_path(v)
|
||||
|
||||
|
||||
class ConfigFolderResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
description: str | None
|
||||
mount_path: str
|
||||
files: dict
|
||||
project_overrides: dict | None
|
||||
is_active: bool
|
||||
created_at: str
|
||||
updated_at: str
|
||||
|
||||
|
||||
@router.get("", summary="List config folders", description="Get all config folders for the current user.")
|
||||
async def list_config_folders(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""List config folders for the current user."""
|
||||
query = select(ConfigFolder).where(ConfigFolder.user_id == user_id)
|
||||
result = await session.execute(query)
|
||||
folders = result.scalars().all()
|
||||
|
||||
return {
|
||||
"folders": [
|
||||
{
|
||||
"id": str(f.id),
|
||||
"user_id": str(f.user_id),
|
||||
"name": f.name,
|
||||
"description": f.description,
|
||||
"mount_path": f.mount_path,
|
||||
"files": f.files,
|
||||
"project_overrides": f.project_overrides,
|
||||
"is_active": f.is_active,
|
||||
"created_at": f.created_at.isoformat() if f.created_at else None,
|
||||
"updated_at": f.updated_at.isoformat() if f.updated_at else None,
|
||||
}
|
||||
for f in folders
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@router.post("", summary="Create config folder", description="Create a new config folder.", status_code=status.HTTP_201_CREATED)
|
||||
async def create_config_folder(
|
||||
data: ConfigFolderCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Create a config folder."""
|
||||
# Check for duplicate name
|
||||
existing = await session.scalar(
|
||||
select(ConfigFolder).where(
|
||||
ConfigFolder.user_id == user_id,
|
||||
ConfigFolder.name == data.name,
|
||||
)
|
||||
)
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"config folder with name '{data.name}' already exists"
|
||||
)
|
||||
|
||||
folder = ConfigFolder(
|
||||
user_id=user_id,
|
||||
name=data.name,
|
||||
description=data.description,
|
||||
mount_path=data.mount_path,
|
||||
files=data.files,
|
||||
)
|
||||
session.add(folder)
|
||||
await session.commit()
|
||||
await session.refresh(folder)
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"user_id": str(folder.user_id),
|
||||
"name": folder.name,
|
||||
"description": folder.description,
|
||||
"mount_path": folder.mount_path,
|
||||
"files": folder.files,
|
||||
"project_overrides": folder.project_overrides,
|
||||
"is_active": folder.is_active,
|
||||
"created_at": folder.created_at.isoformat() if folder.created_at else None,
|
||||
"updated_at": folder.updated_at.isoformat() if folder.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{folder_id}", summary="Update config folder", description="Update an existing config folder.")
|
||||
async def update_config_folder(
|
||||
folder_id: uuid.UUID,
|
||||
data: ConfigFolderUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Update a config folder."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
if data.name is not None:
|
||||
folder.name = data.name
|
||||
if data.description is not None:
|
||||
folder.description = data.description
|
||||
if data.mount_path is not None:
|
||||
folder.mount_path = data.mount_path
|
||||
if data.files is not None:
|
||||
folder.files = data.files
|
||||
if data.is_active is not None:
|
||||
folder.is_active = data.is_active
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(folder)
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"user_id": str(folder.user_id),
|
||||
"name": folder.name,
|
||||
"description": folder.description,
|
||||
"mount_path": folder.mount_path,
|
||||
"files": folder.files,
|
||||
"project_overrides": folder.project_overrides,
|
||||
"is_active": folder.is_active,
|
||||
"created_at": folder.created_at.isoformat() if folder.created_at else None,
|
||||
"updated_at": folder.updated_at.isoformat() if folder.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{folder_id}", summary="Delete config folder", description="Delete a config folder.", status_code=status.HTTP_204_NO_CONTENT)
|
||||
async def delete_config_folder(
|
||||
folder_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Delete a config folder."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
await session.delete(folder)
|
||||
await session.commit()
|
||||
|
||||
|
||||
class ProjectOverrideWithId(ProjectOverrideCreate):
|
||||
project_id: uuid.UUID = Field(description="Project ID for the override")
|
||||
|
||||
|
||||
@router.get("/{folder_id}", summary="Get config folder by ID", description="Get a single config folder by its ID.")
|
||||
async def get_config_folder(
|
||||
folder_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Get a config folder by ID."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"user_id": str(folder.user_id),
|
||||
"name": folder.name,
|
||||
"description": folder.description,
|
||||
"mount_path": folder.mount_path,
|
||||
"files": folder.files,
|
||||
"project_overrides": folder.project_overrides,
|
||||
"is_active": folder.is_active,
|
||||
"created_at": folder.created_at.isoformat() if folder.created_at else None,
|
||||
"updated_at": folder.updated_at.isoformat() if folder.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/{folder_id}/overrides", summary="Add project override", description="Add a project override to a config folder.")
|
||||
async def add_project_override(
|
||||
folder_id: uuid.UUID,
|
||||
data: ProjectOverrideWithId,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Add a project override to a config folder."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
# Initialize project_overrides if None
|
||||
if folder.project_overrides is None:
|
||||
folder.project_overrides = {}
|
||||
|
||||
# Add/update override
|
||||
override_data = {}
|
||||
if data.mount_path is not None:
|
||||
override_data["mount_path"] = data.mount_path
|
||||
if data.files is not None:
|
||||
override_data["files"] = data.files
|
||||
|
||||
# Use a copy to trigger SQLAlchemy change detection on JSONB
|
||||
current_overrides = dict(folder.project_overrides or {})
|
||||
current_overrides[str(data.project_id)] = override_data
|
||||
folder.project_overrides = current_overrides
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(folder)
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"project_overrides": folder.project_overrides,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{folder_id}/overrides/{project_id}", summary="Update project override", description="Update a project override.")
|
||||
async def update_project_override(
|
||||
folder_id: uuid.UUID,
|
||||
project_id: uuid.UUID,
|
||||
data: ProjectOverrideCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Update a project override."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
# Initialize project_overrides if None
|
||||
if folder.project_overrides is None:
|
||||
folder.project_overrides = {}
|
||||
|
||||
# Update override
|
||||
current_overrides = dict(folder.project_overrides or {})
|
||||
override_data = current_overrides.get(str(project_id), {})
|
||||
if data.mount_path is not None:
|
||||
override_data["mount_path"] = data.mount_path
|
||||
if data.files is not None:
|
||||
override_data["files"] = data.files
|
||||
|
||||
current_overrides[str(project_id)] = override_data
|
||||
folder.project_overrides = current_overrides
|
||||
|
||||
# Mark the field as modified to ensure SQLAlchemy detects the change
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
flag_modified(folder, "project_overrides")
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(folder)
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"project_overrides": folder.project_overrides,
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{folder_id}/overrides/{project_id}", summary="Remove project override", description="Remove a project override.")
|
||||
async def remove_project_override(
|
||||
folder_id: uuid.UUID,
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Remove a project override."""
|
||||
folder = await session.get(ConfigFolder, folder_id)
|
||||
if folder is None or folder.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config folder not found")
|
||||
|
||||
# Remove override if exists
|
||||
current_overrides = dict(folder.project_overrides or {})
|
||||
if str(project_id) in current_overrides:
|
||||
del current_overrides[str(project_id)]
|
||||
folder.project_overrides = current_overrides
|
||||
await session.commit()
|
||||
await session.refresh(folder)
|
||||
|
||||
return {
|
||||
"id": str(folder.id),
|
||||
"project_overrides": folder.project_overrides or {},
|
||||
}
|
||||
@@ -0,0 +1,723 @@
|
||||
"""Config profile API endpoints."""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.api.shared_validators import validate_env_vars as _validate_env_vars
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||
from src.models.project import Project
|
||||
from src.models.tool_type import ToolType
|
||||
from src.services.config_profile_resolver import (
|
||||
ConfigProfileCycleError,
|
||||
check_include_cycle,
|
||||
resolve_profile,
|
||||
resolved_profile_to_dict,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/config-profiles", tags=["config-profiles"])
|
||||
|
||||
MAX_PROFILE_SIZE_MB = 10
|
||||
MAX_PROFILE_SIZE_BYTES = MAX_PROFILE_SIZE_MB * 1024 * 1024
|
||||
|
||||
|
||||
def _validate_uuid(v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
try:
|
||||
uuid.UUID(v)
|
||||
except ValueError:
|
||||
raise ValueError(f"Invalid UUID: {v}")
|
||||
return v
|
||||
|
||||
|
||||
def _calculate_profile_size(data: dict) -> int:
|
||||
"""Calculate approximate serialized size of profile data."""
|
||||
total = 0
|
||||
for key, value in data.get("env_vars", {}).items():
|
||||
total += len(key.encode("utf-8")) + len(str(value).encode("utf-8"))
|
||||
for key, value in data.get("runtime_hints", {}).items():
|
||||
total += len(key.encode("utf-8")) + len(str(value).encode("utf-8"))
|
||||
for mount in data.get("mounts", []):
|
||||
total += len(str(mount.get("target", "")).encode("utf-8"))
|
||||
total += len(str(mount.get("mode", "")).encode("utf-8"))
|
||||
for path, content in mount.get("files", {}).items():
|
||||
total += len(path.encode("utf-8")) + len(content.encode("utf-8"))
|
||||
for path, content in data.get("files", {}).items():
|
||||
total += len(path.encode("utf-8")) + len(content.encode("utf-8"))
|
||||
return total
|
||||
|
||||
|
||||
class GitMountItem(BaseModel):
|
||||
remote_url: str = Field(description="Git remote URL (HTTPS or SSH)")
|
||||
source_path: str = Field(default=".", description="Path within repository (supports glob patterns)")
|
||||
target_path: str = Field(description="Absolute path inside container")
|
||||
branch: str | None = Field(default=None, description="Optional branch or tag name")
|
||||
|
||||
@field_validator("remote_url")
|
||||
@classmethod
|
||||
def validate_remote_url(cls, v: str) -> str:
|
||||
if not v.startswith(("http://", "https://", "git@", "ssh://")):
|
||||
raise ValueError("remote_url must be a valid git URL (https://, git@, or ssh://)")
|
||||
return v
|
||||
|
||||
@field_validator("source_path")
|
||||
@classmethod
|
||||
def validate_source_path(cls, v: str) -> str:
|
||||
if v.startswith("/"):
|
||||
raise ValueError("source_path must be relative (no leading /)")
|
||||
if ".." in v:
|
||||
raise ValueError("source_path cannot contain path traversal (..)")
|
||||
return v
|
||||
|
||||
@field_validator("target_path")
|
||||
@classmethod
|
||||
def validate_target_path(cls, v: str) -> str:
|
||||
if ".." in v:
|
||||
raise ValueError("target_path cannot contain path traversal (..)")
|
||||
return v
|
||||
|
||||
|
||||
class MountItem(BaseModel):
|
||||
target: str = Field(description="Absolute mount target path")
|
||||
mode: str = Field(default="rw", description="Mount mode: ro or rw")
|
||||
files: dict = Field(default_factory=dict, description="Files as {relative_path: content}")
|
||||
|
||||
@field_validator("target")
|
||||
@classmethod
|
||||
def validate_target(cls, v: str) -> str:
|
||||
if not v.startswith("/"):
|
||||
raise ValueError("Mount target must be absolute (start with /)")
|
||||
return v
|
||||
|
||||
@field_validator("mode")
|
||||
@classmethod
|
||||
def validate_mode(cls, v: str) -> str:
|
||||
if v not in ("ro", "rw"):
|
||||
raise ValueError("Mount mode must be 'ro' or 'rw'")
|
||||
return v
|
||||
|
||||
@field_validator("files")
|
||||
@classmethod
|
||||
def validate_files(cls, v: dict) -> dict:
|
||||
for path in v.keys():
|
||||
if ".." in path or not path:
|
||||
raise ValueError(f"Invalid file path: {path}")
|
||||
if path.startswith("/"):
|
||||
raise ValueError(
|
||||
f"Mount file paths must be relative (got: {path}). "
|
||||
f"The mount target defines the absolute container path."
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
class ConfigProfileCreate(BaseModel):
|
||||
name: str = Field(description="Profile name (unique per user)")
|
||||
description: str | None = Field(default=None, description="Optional description")
|
||||
project_id: str | None = Field(default=None, description="Optional project ID")
|
||||
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
||||
env_vars: dict = Field(default_factory=dict, description="Environment variables")
|
||||
runtime_hints: dict = Field(default_factory=dict, description="Runtime hints")
|
||||
mounts: list[MountItem] = Field(default_factory=list, description="Mount definitions")
|
||||
files: dict = Field(default_factory=dict, description="Files as {relative_path: content}")
|
||||
git_mounts: list[GitMountItem] = Field(default_factory=list, description="Git repository mounts")
|
||||
is_default: bool = Field(default=False, description="Whether this is the default profile for its scope")
|
||||
|
||||
@field_validator("project_id", "tool_type_id")
|
||||
@classmethod
|
||||
def validate_uuids(cls, v: str | None) -> str | None:
|
||||
return _validate_uuid(v)
|
||||
|
||||
@field_validator("files")
|
||||
@classmethod
|
||||
def validate_files(cls, v: dict) -> dict:
|
||||
for path in v.keys():
|
||||
if ".." in path or not path:
|
||||
raise ValueError(f"Invalid file path: {path}")
|
||||
if path.startswith("/"):
|
||||
raise ValueError(
|
||||
f"File paths must be relative (got: {path}). "
|
||||
f"Use Mounts for absolute container paths."
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("env_vars")
|
||||
@classmethod
|
||||
def validate_env_vars(cls, v: dict) -> dict:
|
||||
result = _validate_env_vars(v)
|
||||
if result is None:
|
||||
raise ValueError("env_vars must be a JSON object")
|
||||
return result
|
||||
|
||||
@field_validator("runtime_hints")
|
||||
@classmethod
|
||||
def validate_runtime_hints(cls, v: dict) -> dict:
|
||||
if not isinstance(v, dict):
|
||||
raise ValueError("runtime_hints must be a JSON object")
|
||||
return v
|
||||
|
||||
@field_validator("mounts")
|
||||
@classmethod
|
||||
def validate_mounts(cls, v: list) -> list:
|
||||
if not isinstance(v, list):
|
||||
raise ValueError("mounts must be a JSON array")
|
||||
return v
|
||||
|
||||
|
||||
class ConfigProfileUpdate(BaseModel):
|
||||
name: str | None = Field(default=None, description="Profile name")
|
||||
description: str | None = Field(default=None, description="Optional description")
|
||||
project_id: str | None = Field(default=None, description="Optional project ID")
|
||||
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
||||
env_vars: dict | None = Field(default=None, description="Environment variables")
|
||||
runtime_hints: dict | None = Field(default=None, description="Runtime hints")
|
||||
mounts: list[MountItem] | None = Field(default=None, description="Mount definitions")
|
||||
files: dict | None = Field(default=None, description="Files as {relative_path: content}")
|
||||
git_mounts: list[GitMountItem] | None = Field(default=None, description="Git repository mounts")
|
||||
is_default: bool | None = Field(default=None, description="Whether this is the default profile")
|
||||
|
||||
@field_validator("project_id", "tool_type_id")
|
||||
@classmethod
|
||||
def validate_uuids(cls, v: str | None) -> str | None:
|
||||
return _validate_uuid(v)
|
||||
|
||||
@field_validator("files")
|
||||
@classmethod
|
||||
def validate_files(cls, v: dict | None) -> dict | None:
|
||||
if v is None:
|
||||
return v
|
||||
for path in v.keys():
|
||||
if ".." in path or path.startswith("/") or not path:
|
||||
raise ValueError(f"Invalid file path: {path}")
|
||||
return v
|
||||
|
||||
|
||||
class ConfigProfileIncludeUpdate(BaseModel):
|
||||
includes: list[str] = Field(description="Ordered list of included profile IDs")
|
||||
|
||||
@field_validator("includes")
|
||||
@classmethod
|
||||
def validate_includes(cls, v: list) -> list:
|
||||
for item in v:
|
||||
try:
|
||||
uuid.UUID(item)
|
||||
except ValueError:
|
||||
raise ValueError(f"Invalid UUID in includes: {item}")
|
||||
return v
|
||||
|
||||
|
||||
class ConfigProfileResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
description: str | None
|
||||
project_id: str | None
|
||||
tool_type_id: str | None
|
||||
env_vars: dict
|
||||
runtime_hints: dict
|
||||
mounts: list
|
||||
files: dict
|
||||
git_mounts: list
|
||||
is_default: bool
|
||||
includes: list[dict]
|
||||
created_at: str
|
||||
updated_at: str
|
||||
|
||||
|
||||
async def _get_profile_with_includes(session: AsyncSession, profile_id: uuid.UUID) -> ConfigProfile | None:
|
||||
"""Fetch a profile with includes eagerly loaded."""
|
||||
result = await session.execute(
|
||||
select(ConfigProfile)
|
||||
.where(ConfigProfile.id == profile_id)
|
||||
.options(selectinload(ConfigProfile.includes))
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _check_access(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
project_id: uuid.UUID | None = None,
|
||||
tool_type_id: uuid.UUID | None = None,
|
||||
) -> None:
|
||||
"""Verify user has access to referenced project and tool type."""
|
||||
if project_id is not None:
|
||||
project = await session.get(Project, project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
# Add ownership check if needed; for now just verify existence
|
||||
if tool_type_id is not None:
|
||||
tool_type = await session.get(ToolType, tool_type_id)
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Tool type not found")
|
||||
|
||||
|
||||
async def _validate_git_mounts(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
git_mounts: list[dict],
|
||||
project_id: uuid.UUID | None = None,
|
||||
) -> None:
|
||||
"""Validate git mount URLs.
|
||||
|
||||
Simply checks that remote_url looks like a valid git URL.
|
||||
Actual clone validation happens at instance startup time.
|
||||
"""
|
||||
for mount in git_mounts:
|
||||
remote_url = mount.get("remote_url")
|
||||
if not remote_url:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Git mount missing remote_url",
|
||||
)
|
||||
|
||||
if not remote_url.startswith(("http://", "https://", "git@", "ssh://")):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid git URL: {remote_url}",
|
||||
)
|
||||
|
||||
|
||||
def _profile_to_response(profile: ConfigProfile, includes: list[ConfigProfileInclude] | None = None) -> dict:
|
||||
return {
|
||||
"id": str(profile.id),
|
||||
"user_id": str(profile.user_id),
|
||||
"name": profile.name,
|
||||
"description": profile.description,
|
||||
"project_id": str(profile.project_id) if profile.project_id else None,
|
||||
"tool_type_id": str(profile.tool_type_id) if profile.tool_type_id else None,
|
||||
"env_vars": profile.env_vars or {},
|
||||
"runtime_hints": profile.runtime_hints or {},
|
||||
"mounts": profile.mounts or [],
|
||||
"git_mounts": profile.git_mounts or [],
|
||||
"files": profile.files or {},
|
||||
"is_default": profile.is_default,
|
||||
"includes": [
|
||||
{
|
||||
"id": str(inc.id),
|
||||
"included_profile_id": str(inc.included_profile_id),
|
||||
"order_index": inc.order_index,
|
||||
}
|
||||
for inc in (includes or profile.includes)
|
||||
],
|
||||
"created_at": profile.created_at.isoformat() if profile.created_at else None,
|
||||
"updated_at": profile.updated_at.isoformat() if profile.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.get("", response_model=list[ConfigProfileResponse])
|
||||
async def list_config_profiles(
|
||||
project_id: str | None = Query(None, description="Filter by project compatibility"),
|
||||
tool_type_id: str | None = Query(None, description="Filter by tool type compatibility"),
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""List config profiles, optionally filtered by compatibility."""
|
||||
user_uuid = current_user_id
|
||||
query = select(ConfigProfile).where(ConfigProfile.user_id == user_uuid).options(selectinload(ConfigProfile.includes))
|
||||
|
||||
if project_id or tool_type_id:
|
||||
# Compatibility filter: include portable profiles and matching scoped profiles
|
||||
project_uuid = uuid.UUID(project_id) if project_id else None
|
||||
tool_uuid = uuid.UUID(tool_type_id) if tool_type_id else None
|
||||
|
||||
from sqlalchemy import or_
|
||||
|
||||
conditions: list = []
|
||||
# Portable profiles (no project, no tool)
|
||||
conditions.append(
|
||||
(ConfigProfile.project_id.is_(None)) & (ConfigProfile.tool_type_id.is_(None))
|
||||
)
|
||||
if project_uuid:
|
||||
# Profiles matching this project (with or without tool)
|
||||
conditions.append(ConfigProfile.project_id == project_uuid)
|
||||
if tool_uuid:
|
||||
# Profiles matching this tool (with or without project)
|
||||
conditions.append(ConfigProfile.tool_type_id == tool_uuid)
|
||||
if project_uuid and tool_uuid:
|
||||
# Exact match
|
||||
conditions.append(
|
||||
(ConfigProfile.project_id == project_uuid) & (ConfigProfile.tool_type_id == tool_uuid)
|
||||
)
|
||||
|
||||
query = query.where(or_(*conditions))
|
||||
|
||||
result = await session.execute(query)
|
||||
profiles = result.scalars().all()
|
||||
return [_profile_to_response(p) for p in profiles]
|
||||
|
||||
|
||||
@router.post("", response_model=ConfigProfileResponse, status_code=status.HTTP_201_CREATED)
|
||||
async def create_config_profile(
|
||||
data: ConfigProfileCreate,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Create a new config profile."""
|
||||
user_uuid = current_user_id
|
||||
|
||||
# Check for duplicate name
|
||||
existing = await session.execute(
|
||||
select(ConfigProfile).where(
|
||||
ConfigProfile.user_id == user_uuid,
|
||||
ConfigProfile.name == data.name,
|
||||
).options(selectinload(ConfigProfile.includes))
|
||||
)
|
||||
if existing.scalar_one_or_none() is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Profile with name '{data.name}' already exists",
|
||||
)
|
||||
|
||||
# Validate references
|
||||
project_uuid = uuid.UUID(data.project_id) if data.project_id else None
|
||||
tool_uuid = uuid.UUID(data.tool_type_id) if data.tool_type_id else None
|
||||
await _check_access(session, user_uuid, project_uuid, tool_uuid)
|
||||
|
||||
# Validate git mounts reference existing repositories
|
||||
if data.git_mounts:
|
||||
git_mounts_data = [m.model_dump() if hasattr(m, "model_dump") else m for m in data.git_mounts]
|
||||
await _validate_git_mounts(session, user_uuid, git_mounts_data, project_uuid)
|
||||
|
||||
# Check size
|
||||
size = _calculate_profile_size(data.model_dump())
|
||||
if size > MAX_PROFILE_SIZE_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"Profile size exceeds {MAX_PROFILE_SIZE_MB}MB limit",
|
||||
)
|
||||
|
||||
profile = ConfigProfile(
|
||||
user_id=user_uuid,
|
||||
name=data.name,
|
||||
description=data.description,
|
||||
project_id=project_uuid,
|
||||
tool_type_id=tool_uuid,
|
||||
env_vars=data.env_vars,
|
||||
runtime_hints=data.runtime_hints,
|
||||
mounts=[m.model_dump() for m in data.mounts],
|
||||
git_mounts=[m.model_dump() for m in data.git_mounts],
|
||||
files=data.files,
|
||||
is_default=data.is_default,
|
||||
)
|
||||
session.add(profile)
|
||||
await session.commit()
|
||||
|
||||
# Re-fetch with includes to avoid lazy loading issues
|
||||
result = await session.execute(
|
||||
select(ConfigProfile)
|
||||
.where(ConfigProfile.id == profile.id)
|
||||
.options(selectinload(ConfigProfile.includes))
|
||||
)
|
||||
profile = result.scalar_one()
|
||||
|
||||
logger.debug("Created config profile %s for user %s", profile.id, user_uuid)
|
||||
return _profile_to_response(profile)
|
||||
|
||||
|
||||
@router.get("/{profile_id}", response_model=ConfigProfileResponse)
|
||||
async def get_config_profile(
|
||||
profile_id: str,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Get a config profile by ID."""
|
||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
||||
if profile.user_id != current_user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
||||
return _profile_to_response(profile)
|
||||
|
||||
|
||||
@router.put("/{profile_id}", response_model=ConfigProfileResponse)
|
||||
async def update_config_profile(
|
||||
profile_id: str,
|
||||
data: ConfigProfileUpdate,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Update a config profile."""
|
||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
||||
if profile.user_id != current_user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
|
||||
# Handle name uniqueness
|
||||
if "name" in update_data:
|
||||
existing = await session.execute(
|
||||
select(ConfigProfile).where(
|
||||
ConfigProfile.user_id == profile.user_id,
|
||||
ConfigProfile.name == update_data["name"],
|
||||
ConfigProfile.id != profile.id,
|
||||
)
|
||||
)
|
||||
if existing.scalar_one_or_none() is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Profile with name '{update_data['name']}' already exists",
|
||||
)
|
||||
|
||||
# Validate references
|
||||
project_uuid = (
|
||||
uuid.UUID(update_data["project_id"])
|
||||
if "project_id" in update_data and update_data["project_id"]
|
||||
else (profile.project_id if "project_id" not in update_data else None)
|
||||
)
|
||||
tool_uuid = (
|
||||
uuid.UUID(update_data["tool_type_id"])
|
||||
if "tool_type_id" in update_data and update_data["tool_type_id"]
|
||||
else (profile.tool_type_id if "tool_type_id" not in update_data else None)
|
||||
)
|
||||
await _check_access(session, profile.user_id, project_uuid, tool_uuid)
|
||||
|
||||
# Validate git mounts reference existing repositories
|
||||
if "git_mounts" in update_data and update_data["git_mounts"] is not None:
|
||||
git_mounts_data = [
|
||||
m.model_dump() if hasattr(m, "model_dump") else m
|
||||
for m in update_data["git_mounts"]
|
||||
]
|
||||
await _validate_git_mounts(session, profile.user_id, git_mounts_data, project_uuid)
|
||||
|
||||
# Check size
|
||||
current_data = _profile_to_response(profile)
|
||||
merged = {**current_data, **update_data}
|
||||
size = _calculate_profile_size(merged)
|
||||
if size > MAX_PROFILE_SIZE_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"Profile size exceeds {MAX_PROFILE_SIZE_MB}MB limit",
|
||||
)
|
||||
|
||||
# Apply updates
|
||||
for field_name, value in update_data.items():
|
||||
if field_name in ("project_id", "tool_type_id"):
|
||||
value = uuid.UUID(value) if value else None
|
||||
elif field_name == "mounts" and value is not None:
|
||||
value = [m.model_dump() if not isinstance(m, dict) else m for m in value]
|
||||
elif field_name == "git_mounts" and value is not None:
|
||||
value = [m.model_dump() if not isinstance(m, dict) else m for m in value]
|
||||
setattr(profile, field_name, value)
|
||||
|
||||
await session.commit()
|
||||
|
||||
# Re-fetch with includes to avoid lazy loading issues
|
||||
result = await session.execute(
|
||||
select(ConfigProfile)
|
||||
.where(ConfigProfile.id == profile.id)
|
||||
.options(selectinload(ConfigProfile.includes))
|
||||
)
|
||||
profile = result.scalar_one()
|
||||
|
||||
logger.debug("Updated config profile %s", profile.id)
|
||||
return _profile_to_response(profile)
|
||||
|
||||
|
||||
@router.delete("/{profile_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
async def delete_config_profile(
|
||||
profile_id: str,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Delete a config profile."""
|
||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
||||
if profile.user_id != current_user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
||||
|
||||
await session.delete(profile)
|
||||
await session.commit()
|
||||
|
||||
logger.debug("Deleted config profile %s", profile_id)
|
||||
return None
|
||||
|
||||
|
||||
@router.put("/{profile_id}/includes", response_model=ConfigProfileResponse)
|
||||
async def update_profile_includes(
|
||||
profile_id: str,
|
||||
data: ConfigProfileIncludeUpdate,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Update the ordered includes for a config profile."""
|
||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
||||
if profile.user_id != current_user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
||||
|
||||
# Validate all included profiles exist and belong to the user
|
||||
included_uuids = [uuid.UUID(inc_id) for inc_id in data.includes]
|
||||
for inc_uuid in included_uuids:
|
||||
inc_profile = await session.get(ConfigProfile, inc_uuid)
|
||||
if inc_profile is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Included profile not found: {inc_uuid}",
|
||||
)
|
||||
if inc_profile.user_id != current_user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"Not authorized to include profile: {inc_uuid}",
|
||||
)
|
||||
if inc_uuid == profile.id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Profile cannot include itself",
|
||||
)
|
||||
|
||||
# Check for cycles
|
||||
cycle = await check_include_cycle(session, profile.id, None)
|
||||
if cycle is None and included_uuids:
|
||||
# Check each new include would not create a cycle
|
||||
for inc_uuid in included_uuids:
|
||||
cycle = await check_include_cycle(session, profile.id, inc_uuid)
|
||||
if cycle is not None:
|
||||
break
|
||||
|
||||
if cycle is not None:
|
||||
cycle_str = " -> ".join(str(c) for c in cycle)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Include cycle detected: {cycle_str}",
|
||||
)
|
||||
|
||||
# Remove existing includes
|
||||
result = await session.execute(
|
||||
select(ConfigProfileInclude).where(ConfigProfileInclude.profile_id == profile.id)
|
||||
)
|
||||
for existing in result.scalars().all():
|
||||
await session.delete(existing)
|
||||
await session.flush()
|
||||
|
||||
# Add new includes
|
||||
for order_index, inc_uuid in enumerate(included_uuids):
|
||||
include = ConfigProfileInclude(
|
||||
profile_id=profile.id,
|
||||
included_profile_id=inc_uuid,
|
||||
order_index=order_index,
|
||||
)
|
||||
session.add(include)
|
||||
await session.flush()
|
||||
|
||||
await session.commit()
|
||||
|
||||
# Re-fetch profile (includes loaded separately due to SQLite async issue)
|
||||
result = await session.execute(
|
||||
select(ConfigProfile).where(ConfigProfile.id == profile.id)
|
||||
)
|
||||
profile = result.scalar_one()
|
||||
|
||||
inc_result = await session.execute(
|
||||
select(ConfigProfileInclude).where(ConfigProfileInclude.profile_id == profile.id)
|
||||
)
|
||||
direct_includes = inc_result.scalars().all()
|
||||
|
||||
logger.debug("Updated includes for config profile %s", profile.id)
|
||||
return _profile_to_response(profile, list(direct_includes))
|
||||
|
||||
|
||||
@router.get("/{profile_id}/preview")
|
||||
async def preview_config_profile(
|
||||
profile_id: str,
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Preview the resolved output of a config profile."""
|
||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
||||
if profile.user_id != current_user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
||||
|
||||
try:
|
||||
resolved = await resolve_profile(session, profile.id)
|
||||
except ConfigProfileCycleError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc),
|
||||
)
|
||||
|
||||
return resolved_profile_to_dict(resolved)
|
||||
|
||||
|
||||
@router.get("/defaults/resolve")
|
||||
async def resolve_default_profile(
|
||||
project_id: str = Query(..., description="Project ID"),
|
||||
tool_type_id: str = Query(..., description="Tool type ID"),
|
||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
):
|
||||
"""Resolve the default config profile for a project/tool combination.
|
||||
|
||||
Selects by specificity:
|
||||
1. project+tool explicit default
|
||||
2. project explicit default
|
||||
3. tool explicit default
|
||||
4. global/user explicit default
|
||||
5. first created compatible profile
|
||||
6. none (returns null)
|
||||
"""
|
||||
user_uuid = current_user_id
|
||||
project_uuid = uuid.UUID(project_id)
|
||||
tool_uuid = uuid.UUID(tool_type_id)
|
||||
|
||||
# Fetch all compatible profiles ordered by created_at
|
||||
query = (
|
||||
select(ConfigProfile)
|
||||
.where(ConfigProfile.user_id == user_uuid)
|
||||
.where(
|
||||
(ConfigProfile.project_id.is_(None) & ConfigProfile.tool_type_id.is_(None))
|
||||
| (ConfigProfile.project_id == project_uuid)
|
||||
| (ConfigProfile.tool_type_id == tool_uuid)
|
||||
| (
|
||||
(ConfigProfile.project_id == project_uuid)
|
||||
& (ConfigProfile.tool_type_id == tool_uuid)
|
||||
)
|
||||
)
|
||||
.order_by(ConfigProfile.created_at)
|
||||
)
|
||||
result = await session.execute(query)
|
||||
profiles = result.scalars().all()
|
||||
|
||||
if not profiles:
|
||||
return {"profile_id": None, "profile_name": None}
|
||||
|
||||
# Check explicit defaults by specificity
|
||||
explicit_defaults = [p for p in profiles if p.is_default]
|
||||
|
||||
# Most specific: project+tool
|
||||
for p in explicit_defaults:
|
||||
if p.project_id == project_uuid and p.tool_type_id == tool_uuid:
|
||||
return {"profile_id": str(p.id), "profile_name": p.name}
|
||||
|
||||
# Next: project only
|
||||
for p in explicit_defaults:
|
||||
if p.project_id == project_uuid and p.tool_type_id is None:
|
||||
return {"profile_id": str(p.id), "profile_name": p.name}
|
||||
|
||||
# Next: tool only
|
||||
for p in explicit_defaults:
|
||||
if p.project_id is None and p.tool_type_id == tool_uuid:
|
||||
return {"profile_id": str(p.id), "profile_name": p.name}
|
||||
|
||||
# Next: global/user (no project, no tool)
|
||||
for p in explicit_defaults:
|
||||
if p.project_id is None and p.tool_type_id is None:
|
||||
return {"profile_id": str(p.id), "profile_name": p.name}
|
||||
|
||||
# Fall back to first created compatible profile
|
||||
first = profiles[0]
|
||||
return {"profile_id": str(first.id), "profile_name": first.name}
|
||||
@@ -0,0 +1,65 @@
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.ssh_key import SSHKey
|
||||
|
||||
router = APIRouter(prefix="/dashboard", tags=["dashboard"])
|
||||
|
||||
|
||||
@router.get(
|
||||
"/summary",
|
||||
summary="Get dashboard summary",
|
||||
description="Get a summary of the user's projects, repositories, SSH keys, and recent activity.",
|
||||
)
|
||||
async def get_dashboard_summary(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Get a summary of the user's dashboard data.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with counts of projects, repositories, SSH keys, and recent activity.
|
||||
"""
|
||||
# Count user's projects
|
||||
projects_result = await session.execute(
|
||||
select(func.count()).select_from(Project).where(Project.owner_id == user_id)
|
||||
)
|
||||
projects_count = projects_result.scalar() or 0
|
||||
|
||||
# Count user's repositories
|
||||
repos_result = await session.execute(
|
||||
select(func.count()).select_from(GitRepository).where(GitRepository.owner_id == user_id)
|
||||
)
|
||||
repos_count = repos_result.scalar() or 0
|
||||
|
||||
# Count user's SSH keys
|
||||
ssh_keys_result = await session.execute(
|
||||
select(func.count()).select_from(SSHKey).where(SSHKey.user_id == user_id)
|
||||
)
|
||||
ssh_keys_count = ssh_keys_result.scalar() or 0
|
||||
|
||||
# Get recent activity (latest 5 projects)
|
||||
recent_projects = await session.execute(
|
||||
select(Project)
|
||||
.where(Project.owner_id == user_id)
|
||||
.order_by(Project.created_at.desc())
|
||||
.limit(5)
|
||||
)
|
||||
recent_activity = [f"Created project: {p.name}" for p in recent_projects.scalars().all()]
|
||||
|
||||
return {
|
||||
"projects": projects_count,
|
||||
"repositories": repos_count,
|
||||
"sshKeys": ssh_keys_count,
|
||||
"recentActivity": recent_activity,
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,148 @@
|
||||
"""Health check endpoints and models."""
|
||||
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import text
|
||||
|
||||
from src.database import SessionLocal
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# Track start time for uptime
|
||||
_start_time = time.time()
|
||||
|
||||
|
||||
class DatabaseHealth(BaseModel):
|
||||
"""Database health check result."""
|
||||
|
||||
status: str = Field(description="Database health status", examples=["healthy"])
|
||||
response_time_ms: float = Field(description="Query response time in milliseconds", examples=[5.2])
|
||||
|
||||
|
||||
class DiskHealth(BaseModel):
|
||||
"""Disk space health check result."""
|
||||
|
||||
status: str = Field(description="Disk health status", examples=["healthy"])
|
||||
free_gb: float = Field(description="Free disk space in GB", examples=[45.2])
|
||||
total_gb: float = Field(description="Total disk space in GB", examples=[100.0])
|
||||
|
||||
|
||||
class HealthChecks(BaseModel):
|
||||
"""Individual health checks."""
|
||||
|
||||
database: DatabaseHealth | None = None
|
||||
disk: DiskHealth | None = None
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""Overall health check response."""
|
||||
|
||||
status: str = Field(description="Overall health status", examples=["healthy"])
|
||||
timestamp: str = Field(description="ISO 8601 timestamp", examples=["2026-05-19T12:00:00Z"])
|
||||
version: str = Field(description="API version", examples=["0.1.0"])
|
||||
checks: HealthChecks = Field(description="Individual health checks")
|
||||
uptime_seconds: float = Field(description="Server uptime in seconds", examples=[3600.0])
|
||||
|
||||
|
||||
class DatabaseHealthResponse(BaseModel):
|
||||
"""Database-specific health check response."""
|
||||
|
||||
status: str = Field(description="Database health status", examples=["healthy"])
|
||||
response_time_ms: float = Field(description="Query response time in milliseconds", examples=[5.2])
|
||||
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
response_model=HealthResponse,
|
||||
summary="Health check",
|
||||
description="Returns overall system health status including database and disk checks.",
|
||||
tags=["Health"],
|
||||
)
|
||||
async def health_check() -> dict[str, Any]:
|
||||
"""Check overall system health.
|
||||
|
||||
Returns:
|
||||
HealthResponse with status, timestamp, version, checks, and uptime.
|
||||
"""
|
||||
checks = HealthChecks()
|
||||
overall_status = "healthy"
|
||||
|
||||
# Database check
|
||||
try:
|
||||
import time as time_module
|
||||
|
||||
start = time_module.perf_counter()
|
||||
async with SessionLocal() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
db_time = (time_module.perf_counter() - start) * 1000
|
||||
checks.database = DatabaseHealth(
|
||||
status="healthy",
|
||||
response_time_ms=round(db_time, 2),
|
||||
)
|
||||
except Exception:
|
||||
checks.database = DatabaseHealth(
|
||||
status="unhealthy",
|
||||
response_time_ms=0.0,
|
||||
)
|
||||
overall_status = "degraded"
|
||||
|
||||
# Disk check
|
||||
try:
|
||||
import shutil
|
||||
|
||||
disk = shutil.disk_usage("/")
|
||||
free_gb = disk.free / (1024**3)
|
||||
total_gb = disk.total / (1024**3)
|
||||
disk_status = "healthy" if free_gb > 1.0 else "degraded"
|
||||
if disk_status == "degraded":
|
||||
overall_status = "degraded"
|
||||
checks.disk = DiskHealth(
|
||||
status=disk_status,
|
||||
free_gb=round(free_gb, 2),
|
||||
total_gb=round(total_gb, 2),
|
||||
)
|
||||
except Exception:
|
||||
checks.disk = None
|
||||
|
||||
return HealthResponse(
|
||||
status=overall_status,
|
||||
timestamp=datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
|
||||
version="0.1.0",
|
||||
checks=checks,
|
||||
uptime_seconds=round(time.time() - _start_time, 2),
|
||||
).model_dump()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/health/db",
|
||||
response_model=DatabaseHealthResponse,
|
||||
summary="Database health check",
|
||||
description="Returns database-specific health status with response time.",
|
||||
tags=["Health"],
|
||||
)
|
||||
async def health_check_db() -> dict[str, Any]:
|
||||
"""Check database health.
|
||||
|
||||
Returns:
|
||||
DatabaseHealthResponse with status and response time.
|
||||
"""
|
||||
import time as time_module
|
||||
|
||||
try:
|
||||
start = time_module.perf_counter()
|
||||
async with SessionLocal() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
db_time = (time_module.perf_counter() - start) * 1000
|
||||
return DatabaseHealthResponse(
|
||||
status="healthy",
|
||||
response_time_ms=round(db_time, 2),
|
||||
).model_dump()
|
||||
except Exception:
|
||||
return DatabaseHealthResponse(
|
||||
status="unhealthy",
|
||||
response_time_ms=0.0,
|
||||
).model_dump()
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Instance proxy router for forwarding HTTP requests to running containers."""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.tool_instance import ToolInstance
|
||||
from src.models.tool_type import ToolType
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/instances", tags=["instance-proxy"])
|
||||
|
||||
|
||||
async def _proxy_request(
|
||||
request: Request,
|
||||
instance_id: uuid.UUID,
|
||||
path: str,
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
) -> Response:
|
||||
"""Proxy an HTTP request to a running instance."""
|
||||
instance = await session.get(ToolInstance, instance_id)
|
||||
if instance is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="instance not found"
|
||||
)
|
||||
|
||||
# Verify ownership
|
||||
if instance.owner_id != user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="not authorized to access this instance",
|
||||
)
|
||||
|
||||
if instance.status != "running" or not instance.container_name:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="instance is not running",
|
||||
)
|
||||
|
||||
# Get the tool type to find the internal port
|
||||
tool_type = await session.get(ToolType, instance.tool_type_id)
|
||||
internal_port = tool_type.default_port if tool_type and tool_type.default_port else instance.port
|
||||
|
||||
# Build target URL using internal port
|
||||
target_url = f"http://{instance.container_name}:{internal_port}"
|
||||
if path:
|
||||
target_url += f"/{path}"
|
||||
|
||||
# Get query string
|
||||
query_string = str(request.query_params)
|
||||
if query_string:
|
||||
target_url += f"?{query_string}"
|
||||
|
||||
# Forward headers (excluding host and cookies)
|
||||
headers: dict[str, str] = {}
|
||||
for key, value in request.headers.items():
|
||||
if key.lower() not in ("host", "cookie", "content-length"):
|
||||
headers[key] = value
|
||||
|
||||
# Forward the request
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
body = await request.body()
|
||||
response = await client.request(
|
||||
method=request.method,
|
||||
url=target_url,
|
||||
headers=headers,
|
||||
content=body,
|
||||
follow_redirects=False,
|
||||
timeout=30.0,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("Proxy error to %s: %s", target_url, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"failed to reach instance: {exc}",
|
||||
)
|
||||
|
||||
# Build response
|
||||
response_headers = dict(response.headers)
|
||||
# Remove hop-by-hop headers
|
||||
for header in ("content-encoding", "transfer-encoding", "connection"):
|
||||
response_headers.pop(header, None)
|
||||
|
||||
return Response(
|
||||
content=response.content,
|
||||
status_code=response.status_code,
|
||||
headers=response_headers,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{instance_id}/proxy/{path:path}")
|
||||
@router.post("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
@router.put("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
@router.delete("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
@router.patch("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
@router.head("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
@router.options("/{instance_id}/proxy/{path:path}", include_in_schema=False)
|
||||
async def proxy_to_instance(
|
||||
request: Request,
|
||||
instance_id: uuid.UUID,
|
||||
path: str = "",
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Response:
|
||||
"""Proxy requests to a running tool instance.
|
||||
|
||||
Args:
|
||||
request: The incoming HTTP request.
|
||||
instance_id: UUID of the instance.
|
||||
path: The path to proxy to the instance.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Response from the proxied instance.
|
||||
"""
|
||||
return await _proxy_request(request, instance_id, path, user_id, session)
|
||||
+105
-45
@@ -1,49 +1,20 @@
|
||||
import os
|
||||
import shutil
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, Response, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.jwt_service import decode_access_token
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.auth.dependencies import _get_owned_project, _get_user, get_current_user_id, get_db_session
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/projects", tags=["projects"])
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
access_token: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
try:
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return uuid.UUID(str(claims["sub"]))
|
||||
except Exception:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid access token")
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
class ProjectCreate(BaseModel):
|
||||
name: str
|
||||
description: str | None = None
|
||||
@@ -68,12 +39,28 @@ class SetDefaultSSHKeyRequest(BaseModel):
|
||||
ssh_key_id: uuid.UUID
|
||||
|
||||
|
||||
@router.post("", response_model=ProjectResponse, status_code=status.HTTP_201_CREATED)
|
||||
@router.post(
|
||||
"",
|
||||
response_model=ProjectResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Create a new project",
|
||||
description="Create a new project for the authenticated user.",
|
||||
)
|
||||
async def create_project(
|
||||
data: ProjectCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
"""Create a new project.
|
||||
|
||||
Args:
|
||||
data: Project creation data including name and optional description.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The newly created project.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
project = Project(
|
||||
name=data.name,
|
||||
@@ -87,36 +74,78 @@ async def create_project(
|
||||
return project
|
||||
|
||||
|
||||
@router.get("", response_model=list[ProjectResponse])
|
||||
@router.get(
|
||||
"",
|
||||
response_model=list[ProjectResponse],
|
||||
summary="List all projects",
|
||||
description="Retrieve all projects owned by the authenticated user.",
|
||||
)
|
||||
async def list_projects(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list[Project]:
|
||||
"""List all projects for the authenticated user.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
List of projects owned by the user.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
result = await session.execute(select(Project).where(Project.owner_id == user.id))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _get_owned_project(
|
||||
@router.get(
|
||||
"/{project_id}",
|
||||
response_model=ProjectResponse,
|
||||
summary="Get a project",
|
||||
description="Retrieve a specific project by ID.",
|
||||
)
|
||||
async def get_project(
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
project = await session.get(Project, project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="project not found")
|
||||
if project.owner_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="not project owner")
|
||||
return project
|
||||
"""Get a specific project by ID.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project to retrieve.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The requested project.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
return await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
|
||||
@router.patch("/{project_id}", response_model=ProjectResponse)
|
||||
@router.patch(
|
||||
"/{project_id}",
|
||||
response_model=ProjectResponse,
|
||||
summary="Update a project",
|
||||
description="Update a project's name or description.",
|
||||
)
|
||||
async def update_project(
|
||||
project_id: uuid.UUID,
|
||||
data: ProjectUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
"""Update a project.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project to update.
|
||||
data: Project update data with optional name and description.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated project.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
@@ -130,12 +159,27 @@ async def update_project(
|
||||
return project
|
||||
|
||||
|
||||
@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@router.delete(
|
||||
"/{project_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="Delete a project",
|
||||
description="Delete a project and all its associated repositories.",
|
||||
)
|
||||
async def delete_project(
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Response:
|
||||
"""Delete a project and all its repositories.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project to delete.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Empty response with 204 status code.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
@@ -152,13 +196,29 @@ async def delete_project(
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.patch("/{project_id}/default-ssh-key", response_model=ProjectResponse)
|
||||
@router.patch(
|
||||
"/{project_id}/default-ssh-key",
|
||||
response_model=ProjectResponse,
|
||||
summary="Set default SSH key",
|
||||
description="Set the default SSH key for a project.",
|
||||
)
|
||||
async def set_default_ssh_key(
|
||||
project_id: uuid.UUID,
|
||||
data: SetDefaultSSHKeyRequest,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
"""Set the default SSH key for a project.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project.
|
||||
data: Request containing the SSH key ID to set as default.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated project.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Shared Pydantic validators for API schemas."""
|
||||
|
||||
|
||||
MAX_FOLDER_SIZE_MB = 10
|
||||
MAX_FOLDER_SIZE_BYTES = MAX_FOLDER_SIZE_MB * 1024 * 1024
|
||||
|
||||
|
||||
def validate_mount_path(v: str | None) -> str | None:
|
||||
"""Validate that a mount path is absolute (starts with /).
|
||||
|
||||
Args:
|
||||
v: Mount path string or None.
|
||||
|
||||
Returns:
|
||||
The validated path, or None if input was None.
|
||||
|
||||
Raises:
|
||||
ValueError: If path is not absolute.
|
||||
"""
|
||||
if v is None:
|
||||
return v
|
||||
if not v.startswith("/"):
|
||||
raise ValueError("Mount path must be absolute (start with /)")
|
||||
return v
|
||||
|
||||
|
||||
def validate_files(v: dict | None, max_size_bytes: int = MAX_FOLDER_SIZE_BYTES) -> dict | None:
|
||||
"""Validate file dict for path traversal and size limits.
|
||||
|
||||
Args:
|
||||
v: Dict of {path: content} or None.
|
||||
max_size_bytes: Maximum total size in bytes.
|
||||
|
||||
Returns:
|
||||
The validated dict, or None if input was None.
|
||||
|
||||
Raises:
|
||||
ValueError: If path traversal detected or size limit exceeded.
|
||||
"""
|
||||
if v is None:
|
||||
return v
|
||||
|
||||
total_size = 0
|
||||
for path, content in v.items():
|
||||
# Check for path traversal
|
||||
if ".." in path or path.startswith("/"):
|
||||
raise ValueError(f"Invalid file path: {path}")
|
||||
total_size += len(content.encode("utf-8"))
|
||||
|
||||
if total_size > max_size_bytes:
|
||||
raise ValueError(f"Total folder size exceeds {max_size_bytes // (1024 * 1024)}MB limit")
|
||||
|
||||
return v
|
||||
|
||||
|
||||
def validate_env_vars(v: dict | None) -> dict | None:
|
||||
"""Validate that environment variables is a JSON object.
|
||||
|
||||
Args:
|
||||
v: Dict of env vars or None.
|
||||
|
||||
Returns:
|
||||
The validated dict, or None if input was None.
|
||||
|
||||
Raises:
|
||||
ValueError: If not a dict.
|
||||
"""
|
||||
if v is None:
|
||||
return v
|
||||
if not isinstance(v, dict):
|
||||
raise ValueError("environment_variables must be a JSON object")
|
||||
return v
|
||||
|
||||
|
||||
def validate_volumes(v: list | None) -> list | None:
|
||||
"""Validate volume mounts list.
|
||||
|
||||
Args:
|
||||
v: List of volume dicts or None.
|
||||
|
||||
Returns:
|
||||
The validated list, or None if input was None.
|
||||
|
||||
Raises:
|
||||
ValueError: If not a list or missing required fields.
|
||||
"""
|
||||
if v is None:
|
||||
return v
|
||||
if not isinstance(v, list):
|
||||
raise ValueError("volumes must be a JSON array")
|
||||
for i, vol in enumerate(v):
|
||||
if not isinstance(vol, dict):
|
||||
raise ValueError(f"Volume at index {i} must be an object")
|
||||
if "source" not in vol:
|
||||
raise ValueError(f"Volume at index {i} must have 'source' field")
|
||||
if "target" not in vol:
|
||||
raise ValueError(f"Volume at index {i} must have 'target' field")
|
||||
return v
|
||||
+159
-35
@@ -1,56 +1,41 @@
|
||||
import base64
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.jwt_service import decode_access_token
|
||||
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/ssh-keys", tags=["ssh-keys"])
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
access_token: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
try:
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return uuid.UUID(str(claims["sub"]))
|
||||
except Exception:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid access token")
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
def _get_fernet() -> Fernet:
|
||||
"""Generate a valid Fernet key from the session secret."""
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
settings = Settings()
|
||||
key = settings.jwt_secret[:32].ljust(32, "=")
|
||||
return Fernet(key.encode())
|
||||
# Derive a 32-byte key from the session secret using SHA256
|
||||
key_bytes = hashlib.sha256(settings.session_secret.encode()).digest()
|
||||
# Base64 encode it for Fernet (must be 32 url-safe base64-encoded bytes)
|
||||
key = base64.urlsafe_b64encode(key_bytes)
|
||||
return Fernet(key)
|
||||
|
||||
|
||||
def generate_ssh_key_pair() -> tuple[str, str]:
|
||||
"""Generate a new Ed25519 SSH key pair.
|
||||
|
||||
Returns:
|
||||
Tuple of (private_key, public_key) as strings.
|
||||
"""
|
||||
private_key = Ed25519PrivateKey.generate()
|
||||
public_key = private_key.public_key()
|
||||
|
||||
@@ -81,12 +66,45 @@ class SSHKeyResponse(BaseModel):
|
||||
created_at: datetime
|
||||
|
||||
|
||||
@router.post("", response_model=SSHKeyResponse, status_code=status.HTTP_201_CREATED)
|
||||
class SignPayloadRequest(BaseModel):
|
||||
payload: str
|
||||
|
||||
|
||||
class SignatureResponse(BaseModel):
|
||||
signature: str
|
||||
|
||||
|
||||
class VerifySignatureRequest(BaseModel):
|
||||
payload: str
|
||||
signature: str
|
||||
|
||||
|
||||
class VerifySignatureResponse(BaseModel):
|
||||
valid: bool
|
||||
|
||||
|
||||
@router.post(
|
||||
"",
|
||||
response_model=SSHKeyResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Create SSH key",
|
||||
description="Generate a new Ed25519 SSH key pair for the authenticated user.",
|
||||
)
|
||||
async def create_ssh_key(
|
||||
data: SSHKeyCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> SSHKey:
|
||||
"""Create a new SSH key pair.
|
||||
|
||||
Args:
|
||||
data: SSH key creation data including the key name.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The newly created SSH key with public key exposed.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
private_key, public_key = generate_ssh_key_pair()
|
||||
|
||||
@@ -105,22 +123,51 @@ async def create_ssh_key(
|
||||
return ssh_key
|
||||
|
||||
|
||||
@router.get("", response_model=list[SSHKeyResponse])
|
||||
@router.get(
|
||||
"",
|
||||
response_model=list[SSHKeyResponse],
|
||||
summary="List SSH keys",
|
||||
description="List all SSH keys for the authenticated user.",
|
||||
)
|
||||
async def list_ssh_keys(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list[SSHKey]:
|
||||
"""List all SSH keys for the authenticated user.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
List of SSH keys owned by the user.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
result = await session.execute(select(SSHKey).where(SSHKey.user_id == user.id))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
@router.delete("/{key_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@router.delete(
|
||||
"/{key_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="Delete SSH key",
|
||||
description="Delete an SSH key by ID.",
|
||||
)
|
||||
async def delete_ssh_key(
|
||||
key_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Delete an SSH key.
|
||||
|
||||
Args:
|
||||
key_id: UUID of the SSH key to delete.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
None with 204 status code.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
ssh_key = await session.get(SSHKey, key_id)
|
||||
if ssh_key is None or ssh_key.user_id != user.id:
|
||||
@@ -128,3 +175,80 @@ async def delete_ssh_key(
|
||||
|
||||
await session.delete(ssh_key)
|
||||
await session.commit()
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{key_id}/sign",
|
||||
response_model=SignatureResponse,
|
||||
summary="Sign payload",
|
||||
description="Sign a payload using the SSH private key.",
|
||||
)
|
||||
async def sign_payload(
|
||||
key_id: uuid.UUID,
|
||||
data: SignPayloadRequest,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> SignatureResponse:
|
||||
"""Sign a payload with an SSH key.
|
||||
|
||||
Args:
|
||||
key_id: UUID of the SSH key to use for signing.
|
||||
data: Sign request containing the payload string.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Base64-encoded Ed25519 signature.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
ssh_key = await session.get(SSHKey, key_id)
|
||||
if ssh_key is None or ssh_key.user_id != user.id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="ssh key not found")
|
||||
|
||||
fernet = _get_fernet()
|
||||
private_key_pem = fernet.decrypt(ssh_key.private_key_encrypted.encode()).decode()
|
||||
|
||||
private_key = serialization.load_ssh_private_key(
|
||||
private_key_pem.encode(), password=None
|
||||
)
|
||||
|
||||
signature = private_key.sign(data.payload.encode())
|
||||
return SignatureResponse(signature=base64.b64encode(signature).decode())
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{key_id}/verify",
|
||||
response_model=VerifySignatureResponse,
|
||||
summary="Verify signature",
|
||||
description="Verify a signature against a payload using the SSH public key.",
|
||||
)
|
||||
async def verify_signature(
|
||||
key_id: uuid.UUID,
|
||||
data: VerifySignatureRequest,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> VerifySignatureResponse:
|
||||
"""Verify a signature with an SSH key's public key.
|
||||
|
||||
Args:
|
||||
key_id: UUID of the SSH key to use for verification.
|
||||
data: Verify request containing payload and base64-encoded signature.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Whether the signature is valid.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
ssh_key = await session.get(SSHKey, key_id)
|
||||
if ssh_key is None or ssh_key.user_id != user.id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="ssh key not found")
|
||||
|
||||
public_key = serialization.load_ssh_public_key(ssh_key.public_key.encode())
|
||||
|
||||
try:
|
||||
signature = base64.b64decode(data.signature)
|
||||
public_key.verify(signature, data.payload.encode())
|
||||
return VerifySignatureResponse(valid=True)
|
||||
except Exception:
|
||||
return VerifySignatureResponse(valid=False)
|
||||
|
||||
@@ -0,0 +1,324 @@
|
||||
"""WebSocket terminal endpoint for tool instances."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.dependencies import get_db_session
|
||||
from src.models.tool_instance import ToolInstance
|
||||
from src.models.tool_type import ToolType
|
||||
from src.services.terminal_manager import terminal_manager
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SessionRef:
|
||||
"""Mutable reference to a terminal session, allowing updates during reset."""
|
||||
|
||||
def __init__(self, session):
|
||||
self.session = session
|
||||
|
||||
|
||||
@router.websocket(
|
||||
"/ws/tool-instances/{instance_id}/terminal",
|
||||
)
|
||||
async def terminal_websocket(
|
||||
websocket: WebSocket,
|
||||
instance_id: str,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""WebSocket endpoint for terminal access to a tool instance.
|
||||
|
||||
Provides an interactive terminal session inside a running tool instance container.
|
||||
Sessions persist across WebSocket disconnections.
|
||||
|
||||
Args:
|
||||
websocket: The WebSocket connection.
|
||||
instance_id: UUID string of the tool instance.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
None. Communicates via WebSocket messages.
|
||||
"""
|
||||
logger.debug("Terminal WebSocket connection attempt for instance %s", instance_id)
|
||||
await websocket.accept()
|
||||
logger.debug("Terminal WebSocket accepted for instance %s", instance_id)
|
||||
|
||||
try:
|
||||
# Parse instance_id
|
||||
instance_uuid = uuid.UUID(instance_id)
|
||||
except ValueError:
|
||||
logger.error("Invalid instance ID: %s", instance_id)
|
||||
await websocket.close(code=4001, reason="Invalid instance ID")
|
||||
return
|
||||
|
||||
# Authenticate user from session cookie
|
||||
user_id = await _get_user_from_websocket(websocket, db_session)
|
||||
if user_id is None:
|
||||
logger.warning("Unauthorized terminal access attempt for instance %s", instance_id)
|
||||
await websocket.close(code=4003, reason="Unauthorized")
|
||||
return
|
||||
|
||||
# Get instance and verify ownership
|
||||
instance = await db_session.get(ToolInstance, instance_uuid)
|
||||
if instance is None:
|
||||
logger.warning("Instance %s not found", instance_id)
|
||||
await websocket.close(code=4004, reason="Instance not found")
|
||||
return
|
||||
|
||||
if instance.owner_id != user_id:
|
||||
logger.warning("Forbidden terminal access for instance %s by user %s", instance_id, user_id)
|
||||
await websocket.close(code=4003, reason="Forbidden")
|
||||
return
|
||||
|
||||
if instance.status != "running" or not instance.container_id:
|
||||
logger.warning("Instance %s not running (status=%s, container_id=%s)", instance_id, instance.status, instance.container_id)
|
||||
await websocket.close(code=4004, reason="Instance not running")
|
||||
return
|
||||
|
||||
logger.debug("Terminal auth passed for instance %s, user %s", instance_id, user_id)
|
||||
|
||||
# Fetch tool type to get startup_command
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
if startup_command:
|
||||
logger.debug("Using startup command for instance %s: %s", instance_id, startup_command)
|
||||
|
||||
# Get or create terminal session
|
||||
try:
|
||||
session = await terminal_manager.get_or_create_session(
|
||||
instance_uuid,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
)
|
||||
logger.debug("Terminal session ready for instance %s (session_id=%s)", instance_id, session.session_id)
|
||||
|
||||
# Attach WebSocket to session
|
||||
await terminal_manager.attach_websocket(session, websocket)
|
||||
logger.debug("WebSocket attached to session for instance %s", instance_id)
|
||||
|
||||
# Send connected status
|
||||
await websocket.send_json({"type": "status", "status": "connected"})
|
||||
logger.debug("Sent connected status for instance %s", instance_id)
|
||||
|
||||
# Use mutable session reference so loops can survive reset
|
||||
session_ref = SessionRef(session)
|
||||
|
||||
# Start I/O loops and heartbeat
|
||||
read_task = asyncio.create_task(_read_loop(session_ref, websocket))
|
||||
write_task = asyncio.create_task(_write_loop(session_ref, websocket, instance_id))
|
||||
heartbeat_task = asyncio.create_task(_heartbeat_loop(websocket))
|
||||
logger.debug("Started terminal loops for instance %s", instance_id)
|
||||
|
||||
# Wait for either task to complete (indicating disconnect or error)
|
||||
done, pending = await asyncio.wait(
|
||||
[read_task, write_task, heartbeat_task],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
|
||||
logger.debug("Terminal loop completed for instance %s, done=%s", instance_id, len(done))
|
||||
|
||||
# Cancel remaining tasks
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Terminal session error for instance %s: %s", instance_id, str(exc), exc_info=True)
|
||||
await websocket.close(code=4000, reason=f"Error: {exc}")
|
||||
finally:
|
||||
# Detach WebSocket, don't kill session
|
||||
try:
|
||||
if 'session' in locals():
|
||||
await terminal_manager.detach_websocket(session, websocket)
|
||||
logger.debug("WebSocket detached from session for instance %s", instance_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _read_loop(session_ref: SessionRef, websocket) -> None:
|
||||
"""Read output from the container and send to WebSocket."""
|
||||
try:
|
||||
while True:
|
||||
session = session_ref.session
|
||||
if not session.is_alive() or session._closed:
|
||||
await asyncio.sleep(0.1)
|
||||
continue
|
||||
data = await session.read_output()
|
||||
if data:
|
||||
try:
|
||||
await websocket.send_bytes(data)
|
||||
except Exception:
|
||||
break
|
||||
else:
|
||||
await asyncio.sleep(0.01)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _write_loop(session_ref: SessionRef, websocket, instance_id: str) -> None:
|
||||
"""Read input from WebSocket and send to container."""
|
||||
try:
|
||||
while True:
|
||||
session = session_ref.session
|
||||
if not session.is_alive() or session._closed:
|
||||
await asyncio.sleep(0.1)
|
||||
continue
|
||||
message = await websocket.receive()
|
||||
if message["type"] == "websocket.receive":
|
||||
if "bytes" in message:
|
||||
await session.write_input(message["bytes"])
|
||||
elif "text" in message:
|
||||
text = message["text"]
|
||||
if text.startswith("{"):
|
||||
# Control message (JSON)
|
||||
import json
|
||||
try:
|
||||
ctrl = json.loads(text)
|
||||
msg_type = ctrl.get("type")
|
||||
|
||||
if msg_type == "resize":
|
||||
cols = ctrl.get("cols", 80)
|
||||
rows = ctrl.get("rows", 24)
|
||||
logger.debug(f"Received resize message for instance {instance_id}: {cols}x{rows}")
|
||||
await session.resize(cols, rows)
|
||||
elif msg_type == "reset":
|
||||
# Reset terminal session
|
||||
logger.debug("Resetting terminal session for instance %s", session.instance_id)
|
||||
await websocket.send_json({"type": "status", "status": "resetting"})
|
||||
|
||||
# Reset the session
|
||||
new_session = await terminal_manager.reset_session(
|
||||
session.instance_id,
|
||||
session.container_id,
|
||||
startup_command=session.startup_command,
|
||||
)
|
||||
|
||||
# Update the mutable session reference so read_loop uses the new session
|
||||
session_ref.session = new_session
|
||||
|
||||
# Attach to new session
|
||||
await terminal_manager.attach_websocket(new_session, websocket)
|
||||
await websocket.send_json({"type": "status", "status": "connected"})
|
||||
|
||||
# Continue the loop with the new session
|
||||
continue
|
||||
|
||||
except json.JSONDecodeError:
|
||||
# Not a valid JSON control message, treat as regular input
|
||||
await session.write_input(text.encode("utf-8"))
|
||||
else:
|
||||
await session.write_input(text.encode("utf-8"))
|
||||
elif message["type"] == "websocket.disconnect":
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _heartbeat_loop(websocket: WebSocket) -> None:
|
||||
"""Send periodic ping messages to detect disconnections."""
|
||||
try:
|
||||
while True:
|
||||
await asyncio.sleep(30) # Ping every 30 seconds
|
||||
try:
|
||||
await websocket.send_json({"type": "ping"})
|
||||
except Exception:
|
||||
# WebSocket is closed or broken
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@router.post(
|
||||
"/projects/{project_id}/repositories/{repo_id}/instances/{instance_id}/terminal/reset",
|
||||
summary="Reset terminal session",
|
||||
description="Reset the terminal session for a tool instance, killing the current shell and starting fresh.",
|
||||
)
|
||||
async def reset_terminal_session(
|
||||
project_id: uuid.UUID,
|
||||
repo_id: uuid.UUID,
|
||||
instance_id: uuid.UUID,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Reset the terminal session for an instance.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project.
|
||||
repo_id: UUID of the repository.
|
||||
instance_id: UUID of the tool instance.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
Dictionary with status message.
|
||||
"""
|
||||
# Get instance and verify it exists and is running
|
||||
instance = await db_session.get(ToolInstance, instance_id)
|
||||
if instance is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="Instance not found"
|
||||
)
|
||||
|
||||
if instance.status != "running" or not instance.container_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Instance is not running"
|
||||
)
|
||||
|
||||
# Fetch tool type to get startup_command
|
||||
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||
startup_command = tool_type.startup_command if tool_type else None
|
||||
|
||||
try:
|
||||
# Reset the session
|
||||
new_session = await terminal_manager.reset_session(
|
||||
instance_id,
|
||||
instance.container_id,
|
||||
startup_command=startup_command,
|
||||
)
|
||||
|
||||
logger.info("Terminal session reset for instance %s (new session_id=%s)", instance_id, new_session.session_id)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": "Terminal session reset successfully",
|
||||
"instance_id": str(instance_id),
|
||||
"session_id": new_session.session_id,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.error("Failed to reset terminal session for instance %s: %s", instance_id, str(exc), exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to reset terminal session: {exc}"
|
||||
)
|
||||
|
||||
|
||||
async def _get_user_from_websocket(
|
||||
websocket: WebSocket,
|
||||
db_session: AsyncSession,
|
||||
) -> uuid.UUID | None:
|
||||
"""Extract and validate user ID from session cookie in WebSocket.
|
||||
|
||||
Args:
|
||||
websocket: The WebSocket connection.
|
||||
db_session: Database session.
|
||||
|
||||
Returns:
|
||||
The user's UUID if authenticated, None otherwise.
|
||||
"""
|
||||
from src.auth.session import decode_session_cookie
|
||||
from src.config import Settings
|
||||
|
||||
session_cookie = websocket.cookies.get("session")
|
||||
if not session_cookie:
|
||||
return None
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
return uuid.UUID(str(payload["user_id"]))
|
||||
except (ValueError, KeyError):
|
||||
return None
|
||||
@@ -0,0 +1,290 @@
|
||||
"""Tool configuration API endpoints."""
|
||||
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.api.shared_validators import validate_env_vars as _validate_env_vars, validate_volumes as _validate_volumes
|
||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||
from src.models.tool_config import ToolConfig
|
||||
from src.models.tool_type import ToolType
|
||||
|
||||
router = APIRouter(prefix="/tool-configs", tags=["tool-configs"])
|
||||
|
||||
|
||||
class ToolConfigCreate(BaseModel):
|
||||
tool_type_id: str = Field(description="UUID of the tool type")
|
||||
project_id: str | None = Field(default=None, description="Optional project ID for project-scoped config")
|
||||
key: str = Field(description="Config key name")
|
||||
value: str = Field(description="Config value")
|
||||
config_type: str = Field(default="env", description="Type: env or file")
|
||||
file_path: str | None = Field(default=None, description="File path for file-type configs")
|
||||
port_override: int | None = Field(default=None, description="Port override (1-65535)")
|
||||
start_command: str | None = Field(default=None, description="Override container start command")
|
||||
working_directory: str | None = Field(default=None, description="Working directory inside container")
|
||||
environment_variables: dict | None = Field(default=None, description="Environment variables as JSON object")
|
||||
volumes: list[dict] | None = Field(default=None, description="Volume mounts as JSON array")
|
||||
|
||||
@field_validator("port_override")
|
||||
@classmethod
|
||||
def validate_port(cls, v: int | None) -> int | None:
|
||||
if v is None:
|
||||
return v
|
||||
if v < 1 or v > 65535:
|
||||
raise ValueError("Port must be between 1 and 65535")
|
||||
return v
|
||||
|
||||
@field_validator("environment_variables")
|
||||
@classmethod
|
||||
def validate_env_vars(cls, v: dict | None) -> dict | None:
|
||||
return _validate_env_vars(v)
|
||||
|
||||
@field_validator("volumes")
|
||||
@classmethod
|
||||
def validate_volumes(cls, v: list | None) -> list | None:
|
||||
return _validate_volumes(v)
|
||||
|
||||
|
||||
class ToolConfigUpdate(BaseModel):
|
||||
key: str | None = Field(default=None, description="Config key name")
|
||||
value: str | None = Field(default=None, description="Config value")
|
||||
config_type: str | None = Field(default=None, description="Type: env or file")
|
||||
file_path: str | None = Field(default=None, description="File path for file-type configs")
|
||||
port_override: int | None = Field(default=None, description="Port override (1-65535)")
|
||||
start_command: str | None = Field(default=None, description="Override container start command")
|
||||
working_directory: str | None = Field(default=None, description="Working directory inside container")
|
||||
environment_variables: dict | None = Field(default=None, description="Environment variables as JSON object")
|
||||
volumes: list[dict] | None = Field(default=None, description="Volume mounts as JSON array")
|
||||
|
||||
@field_validator("port_override")
|
||||
@classmethod
|
||||
def validate_port(cls, v: int | None) -> int | None:
|
||||
if v is None:
|
||||
return v
|
||||
if v < 1 or v > 65535:
|
||||
raise ValueError("Port must be between 1 and 65535")
|
||||
return v
|
||||
|
||||
@field_validator("environment_variables")
|
||||
@classmethod
|
||||
def validate_env_vars(cls, v: dict | None) -> dict | None:
|
||||
return _validate_env_vars(v)
|
||||
|
||||
@field_validator("volumes")
|
||||
@classmethod
|
||||
def validate_volumes(cls, v: list | None) -> list | None:
|
||||
return _validate_volumes(v)
|
||||
|
||||
|
||||
class ToolConfigResponse(BaseModel):
|
||||
id: str
|
||||
tool_type_id: str
|
||||
project_id: str | None
|
||||
key: str
|
||||
value: str
|
||||
config_type: str
|
||||
file_path: str | None
|
||||
port_override: int | None
|
||||
start_command: str | None
|
||||
working_directory: str | None
|
||||
environment_variables: dict | None
|
||||
volumes: list[dict] | None
|
||||
|
||||
|
||||
@router.get("", summary="List tool configs", description="Get all tool configs for the current user.")
|
||||
async def list_configs(
|
||||
tool_type_id: str | None = None,
|
||||
project_id: str | None = None,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list:
|
||||
"""List tool configs for the current user."""
|
||||
query = select(ToolConfig).where(ToolConfig.user_id == user_id)
|
||||
|
||||
if tool_type_id:
|
||||
query = query.where(ToolConfig.tool_type_id == uuid.UUID(tool_type_id))
|
||||
if project_id:
|
||||
query = query.where(ToolConfig.project_id == uuid.UUID(project_id))
|
||||
else:
|
||||
# If no project specified, get only global configs (project_id is None)
|
||||
query = query.where(ToolConfig.project_id.is_(None))
|
||||
|
||||
result = await session.execute(query)
|
||||
configs = result.scalars().all()
|
||||
|
||||
return [
|
||||
{
|
||||
"id": str(c.id),
|
||||
"tool_type_id": str(c.tool_type_id),
|
||||
"project_id": str(c.project_id) if c.project_id else None,
|
||||
"key": c.key,
|
||||
"value": c.value,
|
||||
"config_type": c.config_type,
|
||||
"file_path": c.file_path,
|
||||
"port_override": c.port_override,
|
||||
"start_command": c.start_command,
|
||||
"working_directory": c.working_directory,
|
||||
"environment_variables": c.environment_variables,
|
||||
"volumes": c.volumes,
|
||||
}
|
||||
for c in configs
|
||||
]
|
||||
|
||||
|
||||
@router.post("", summary="Create tool config", description="Create a new tool config.", status_code=status.HTTP_201_CREATED)
|
||||
async def create_config(
|
||||
data: ToolConfigCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Create a tool config."""
|
||||
# Verify tool type exists
|
||||
tool_type = await session.get(ToolType, uuid.UUID(data.tool_type_id))
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
|
||||
# Check for existing config with same key
|
||||
query = select(ToolConfig).where(
|
||||
ToolConfig.user_id == user_id,
|
||||
ToolConfig.tool_type_id == uuid.UUID(data.tool_type_id),
|
||||
ToolConfig.key == data.key,
|
||||
)
|
||||
if data.project_id:
|
||||
query = query.where(ToolConfig.project_id == uuid.UUID(data.project_id))
|
||||
else:
|
||||
query = query.where(ToolConfig.project_id.is_(None))
|
||||
|
||||
existing = await session.scalar(query)
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"config with key '{data.key}' already exists"
|
||||
)
|
||||
|
||||
config = ToolConfig(
|
||||
user_id=user_id,
|
||||
tool_type_id=uuid.UUID(data.tool_type_id),
|
||||
project_id=uuid.UUID(data.project_id) if data.project_id else None,
|
||||
key=data.key,
|
||||
value=data.value,
|
||||
config_type=data.config_type,
|
||||
file_path=data.file_path,
|
||||
port_override=data.port_override,
|
||||
start_command=data.start_command,
|
||||
working_directory=data.working_directory,
|
||||
environment_variables=data.environment_variables,
|
||||
volumes=data.volumes,
|
||||
)
|
||||
session.add(config)
|
||||
await session.commit()
|
||||
await session.refresh(config)
|
||||
|
||||
return {
|
||||
"id": str(config.id),
|
||||
"tool_type_id": str(config.tool_type_id),
|
||||
"project_id": str(config.project_id) if config.project_id else None,
|
||||
"key": config.key,
|
||||
"value": config.value,
|
||||
"config_type": config.config_type,
|
||||
"file_path": config.file_path,
|
||||
"port_override": config.port_override,
|
||||
"start_command": config.start_command,
|
||||
"working_directory": config.working_directory,
|
||||
"environment_variables": config.environment_variables,
|
||||
"volumes": config.volumes,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{config_id}", summary="Update tool config", description="Update an existing tool config.")
|
||||
async def update_config(
|
||||
config_id: uuid.UUID,
|
||||
data: ToolConfigUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Update a tool config."""
|
||||
config = await session.get(ToolConfig, config_id)
|
||||
if config is None or config.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config not found")
|
||||
|
||||
if data.key is not None:
|
||||
config.key = data.key
|
||||
if data.value is not None:
|
||||
config.value = data.value
|
||||
if data.config_type is not None:
|
||||
config.config_type = data.config_type
|
||||
if data.file_path is not None:
|
||||
config.file_path = data.file_path
|
||||
if data.port_override is not None:
|
||||
config.port_override = data.port_override
|
||||
if data.start_command is not None:
|
||||
config.start_command = data.start_command
|
||||
if data.working_directory is not None:
|
||||
config.working_directory = data.working_directory
|
||||
if data.environment_variables is not None:
|
||||
config.environment_variables = data.environment_variables
|
||||
if data.volumes is not None:
|
||||
config.volumes = data.volumes
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(config)
|
||||
|
||||
return {
|
||||
"id": str(config.id),
|
||||
"tool_type_id": str(config.tool_type_id),
|
||||
"project_id": str(config.project_id) if config.project_id else None,
|
||||
"key": config.key,
|
||||
"value": config.value,
|
||||
"config_type": config.config_type,
|
||||
"file_path": config.file_path,
|
||||
"port_override": config.port_override,
|
||||
"start_command": config.start_command,
|
||||
"working_directory": config.working_directory,
|
||||
"environment_variables": config.environment_variables,
|
||||
"volumes": config.volumes,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/defaults/{tool_type_id}", summary="Get default configs", description="Get suggested default configs for a tool type.")
|
||||
async def get_default_configs(
|
||||
tool_type_id: str,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Get suggested default configs for a tool type."""
|
||||
tool_type = await session.get(ToolType, uuid.UUID(tool_type_id))
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
|
||||
# Return suggested defaults based on required_variables
|
||||
defaults = []
|
||||
for var in tool_type.required_variables:
|
||||
defaults.append({
|
||||
"key": var,
|
||||
"value": "",
|
||||
"config_type": "env",
|
||||
"description": f"Required variable: {var}",
|
||||
})
|
||||
|
||||
return {
|
||||
"tool_type_id": tool_type_id,
|
||||
"suggested_configs": defaults,
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{config_id}", summary="Delete tool config", description="Delete a tool config.")
|
||||
async def delete_config(
|
||||
config_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Delete a tool config."""
|
||||
config = await session.get(ToolConfig, config_id)
|
||||
if config is None or config.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="config not found")
|
||||
|
||||
await session.delete(config)
|
||||
await session.commit()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,549 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.api.tool_types_validation import (
|
||||
check_port_exposed,
|
||||
validate_compose_yaml,
|
||||
validate_required_variables,
|
||||
)
|
||||
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/tool-types", tags=["tool-types"])
|
||||
|
||||
|
||||
async def _require_admin(user: User) -> None:
|
||||
"""Check if user has admin privileges.
|
||||
|
||||
For now, all authenticated users can manage tool types.
|
||||
In production, this should check user.role or similar.
|
||||
"""
|
||||
# For now, all authenticated users can manage tool types
|
||||
# In production, check user.role or similar
|
||||
pass
|
||||
|
||||
|
||||
class ToolTypeCreate(BaseModel):
|
||||
name: str
|
||||
display_name: str
|
||||
description: str | None = None
|
||||
default_port: int = 0
|
||||
definition_type: str = "compose"
|
||||
compose_template: str | None = None
|
||||
dockerfile_template: str | None = None
|
||||
build_context: dict | None = None
|
||||
readiness_probe: dict | None = None
|
||||
startup_command: str | None = None
|
||||
required_variables: list[str] = []
|
||||
category: str = "other"
|
||||
interface_type: str = "web"
|
||||
requires_port: bool = True
|
||||
|
||||
@field_validator("definition_type")
|
||||
@classmethod
|
||||
def validate_definition_type(cls, v: str) -> str:
|
||||
if v not in ("compose", "dockerfile"):
|
||||
raise ValueError("definition_type must be 'compose' or 'dockerfile'")
|
||||
return v
|
||||
|
||||
@field_validator("compose_template")
|
||||
@classmethod
|
||||
def validate_compose_template(cls, v: str | None, info) -> str | None:
|
||||
data = info.data
|
||||
if data.get("definition_type") != "compose":
|
||||
return v
|
||||
|
||||
if v is None:
|
||||
raise ValueError("compose_template is required when definition_type is 'compose'")
|
||||
|
||||
validate_compose_yaml(v)
|
||||
return v
|
||||
|
||||
@field_validator("dockerfile_template")
|
||||
@classmethod
|
||||
def validate_dockerfile_template(cls, v: str | None, info) -> str | None:
|
||||
data = info.data
|
||||
if data.get("definition_type") != "dockerfile":
|
||||
return v
|
||||
|
||||
if v is None:
|
||||
raise ValueError("dockerfile_template is required when definition_type is 'dockerfile'")
|
||||
|
||||
if not v.strip().startswith("FROM"):
|
||||
raise ValueError("Dockerfile must start with a FROM instruction")
|
||||
|
||||
return v
|
||||
|
||||
@field_validator("interface_type")
|
||||
@classmethod
|
||||
def validate_interface_type(cls, v: str) -> str:
|
||||
if v not in ("web", "terminal"):
|
||||
raise ValueError("interface_type must be 'web' or 'terminal'")
|
||||
return v
|
||||
|
||||
@field_validator("default_port")
|
||||
@classmethod
|
||||
def validate_default_port(cls, v: int, info) -> int:
|
||||
data = info.data
|
||||
requires_port = data.get("requires_port", True)
|
||||
if not requires_port:
|
||||
return v
|
||||
if v <= 0 or v > 65535:
|
||||
raise ValueError("Port must be between 1 and 65535")
|
||||
return v
|
||||
|
||||
@field_validator("required_variables")
|
||||
@classmethod
|
||||
def validate_required_variables(cls, v: list[str], info) -> list[str]:
|
||||
if not v:
|
||||
return v
|
||||
|
||||
data = info.data
|
||||
if data.get("definition_type") != "compose":
|
||||
return v
|
||||
|
||||
template = data.get("compose_template")
|
||||
if not template:
|
||||
return v
|
||||
|
||||
for var in v:
|
||||
placeholder = f"{{{{{var}}}}}"
|
||||
if placeholder not in template:
|
||||
raise ValueError(f"Required variable '{var}' not found in compose template")
|
||||
|
||||
return v
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_templates(self) -> "ToolTypeCreate":
|
||||
if self.definition_type == "dockerfile" and self.dockerfile_template is None:
|
||||
raise ValueError("dockerfile_template is required when definition_type is 'dockerfile'")
|
||||
if self.definition_type == "compose" and self.compose_template is None:
|
||||
raise ValueError("compose_template is required when definition_type is 'compose'")
|
||||
|
||||
# Validate that default_port is exposed in compose template (only if requires_port)
|
||||
if self.requires_port and self.definition_type == "compose" and self.compose_template:
|
||||
try:
|
||||
parsed = validate_compose_yaml(self.compose_template)
|
||||
except ValueError:
|
||||
return self
|
||||
|
||||
if not check_port_exposed(parsed, self.default_port):
|
||||
raise ValueError(f"Port {self.default_port} is not exposed in the compose template. Add it to the 'ports' section.")
|
||||
|
||||
return self
|
||||
|
||||
|
||||
class ToolTypeUpdate(BaseModel):
|
||||
display_name: str | None = None
|
||||
description: str | None = None
|
||||
default_port: int | None = None
|
||||
definition_type: str | None = None
|
||||
compose_template: str | None = None
|
||||
dockerfile_template: str | None = None
|
||||
build_context: dict | None = None
|
||||
readiness_probe: dict | None = None
|
||||
startup_command: str | None = None
|
||||
required_variables: list[str] | None = None
|
||||
category: str | None = None
|
||||
interface_type: str | None = None
|
||||
requires_port: bool | None = None
|
||||
|
||||
@field_validator("definition_type")
|
||||
@classmethod
|
||||
def validate_definition_type(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
if v not in ("compose", "dockerfile"):
|
||||
raise ValueError("definition_type must be 'compose' or 'dockerfile'")
|
||||
return v
|
||||
|
||||
@field_validator("interface_type")
|
||||
@classmethod
|
||||
def validate_interface_type(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
if v not in ("web", "terminal"):
|
||||
raise ValueError("interface_type must be 'web' or 'terminal'")
|
||||
return v
|
||||
|
||||
@field_validator("compose_template")
|
||||
@classmethod
|
||||
def validate_compose_template(cls, v: str | None, info) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
|
||||
data = info.data
|
||||
definition_type = data.get("definition_type")
|
||||
if definition_type and definition_type != "compose":
|
||||
return v
|
||||
|
||||
validate_compose_yaml(v)
|
||||
return v
|
||||
|
||||
@field_validator("dockerfile_template")
|
||||
@classmethod
|
||||
def validate_dockerfile_template(cls, v: str | None, info) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
|
||||
data = info.data
|
||||
definition_type = data.get("definition_type")
|
||||
if definition_type and definition_type != "dockerfile":
|
||||
return v
|
||||
|
||||
if not v.strip().startswith("FROM"):
|
||||
raise ValueError("Dockerfile must start with a FROM instruction")
|
||||
|
||||
return v
|
||||
|
||||
|
||||
class ToolTypeResponse(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: uuid.UUID
|
||||
name: str
|
||||
display_name: str
|
||||
description: str | None
|
||||
category: str
|
||||
interface_type: str
|
||||
requires_port: bool
|
||||
default_port: int
|
||||
definition_type: str
|
||||
compose_template: str | None
|
||||
dockerfile_template: str | None
|
||||
build_context: dict | None
|
||||
readiness_probe: dict | None
|
||||
startup_command: str | None
|
||||
required_variables: list[str]
|
||||
created_by_id: uuid.UUID | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
@router.post(
|
||||
"",
|
||||
response_model=ToolTypeResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Create tool type",
|
||||
description="Create a new custom tool type with a Docker Compose template.",
|
||||
)
|
||||
async def create_tool_type(
|
||||
data: ToolTypeCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> ToolType:
|
||||
"""Create a new tool type.
|
||||
|
||||
Args:
|
||||
data: Tool type creation data including name, display name, and compose template.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The newly created tool type.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
await _require_admin(user)
|
||||
|
||||
# Check for duplicate name
|
||||
existing = await session.scalar(select(ToolType).where(ToolType.name == data.name))
|
||||
if existing:
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="tool type with this name already exists")
|
||||
|
||||
tool_type = ToolType(
|
||||
name=data.name,
|
||||
display_name=data.display_name,
|
||||
description=data.description,
|
||||
default_port=data.default_port,
|
||||
definition_type=data.definition_type,
|
||||
compose_template=data.compose_template,
|
||||
dockerfile_template=data.dockerfile_template,
|
||||
build_context=data.build_context,
|
||||
readiness_probe=data.readiness_probe,
|
||||
startup_command=data.startup_command,
|
||||
required_variables=data.required_variables,
|
||||
category=data.category,
|
||||
interface_type=data.interface_type,
|
||||
requires_port=data.requires_port,
|
||||
created_by_id=user.id,
|
||||
)
|
||||
session.add(tool_type)
|
||||
await session.commit()
|
||||
await session.refresh(tool_type)
|
||||
return tool_type
|
||||
|
||||
|
||||
@router.get(
|
||||
"",
|
||||
response_model=list[ToolTypeResponse],
|
||||
summary="List tool types",
|
||||
description="List all available tool types including built-in and custom ones.",
|
||||
)
|
||||
async def list_tool_types(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list[ToolType]:
|
||||
"""List all tool types.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
List of all tool types ordered by name.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
result = await session.execute(select(ToolType).order_by(ToolType.name))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{tool_type_id}",
|
||||
response_model=ToolTypeResponse,
|
||||
summary="Get tool type",
|
||||
description="Get a specific tool type by ID.",
|
||||
)
|
||||
async def get_tool_type(
|
||||
tool_type_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> ToolType:
|
||||
"""Get a specific tool type by ID.
|
||||
|
||||
Args:
|
||||
tool_type_id: UUID of the tool type to retrieve.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The requested tool type.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
tool_type = await session.get(ToolType, tool_type_id)
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
return tool_type
|
||||
|
||||
|
||||
@router.put(
|
||||
"/{tool_type_id}",
|
||||
response_model=ToolTypeResponse,
|
||||
summary="Update tool type",
|
||||
description="Update a custom tool type. Built-in tool types cannot be modified.",
|
||||
)
|
||||
async def update_tool_type(
|
||||
tool_type_id: uuid.UUID,
|
||||
data: ToolTypeUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> ToolType:
|
||||
"""Update a tool type.
|
||||
|
||||
Args:
|
||||
tool_type_id: UUID of the tool type to update.
|
||||
data: Tool type update data with optional fields.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated tool type.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
await _require_admin(user)
|
||||
|
||||
tool_type = await session.get(ToolType, tool_type_id)
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
|
||||
# Built-in tool types can now be modified
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
|
||||
# Validate port if being updated
|
||||
requires_port = update_data.get("requires_port", tool_type.requires_port)
|
||||
if "default_port" in update_data and requires_port:
|
||||
new_port = update_data["default_port"]
|
||||
if new_port <= 0 or new_port > 65535:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Port must be between 1 and 65535"
|
||||
)
|
||||
|
||||
# Only validate port exposure for compose definitions
|
||||
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
||||
if definition_type == "compose":
|
||||
template = update_data.get("compose_template", tool_type.compose_template)
|
||||
if template:
|
||||
try:
|
||||
parsed = validate_compose_yaml(template)
|
||||
if not check_port_exposed(parsed, new_port):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Port {new_port} is not exposed in the compose template"
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(e)
|
||||
)
|
||||
|
||||
# Validate required variables for compose definitions
|
||||
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
||||
if definition_type == "compose":
|
||||
if "required_variables" in update_data and "compose_template" in update_data:
|
||||
validate_required_variables(
|
||||
update_data["compose_template"], update_data["required_variables"]
|
||||
)
|
||||
elif "required_variables" in update_data:
|
||||
template = tool_type.compose_template
|
||||
if template:
|
||||
validate_required_variables(template, update_data["required_variables"])
|
||||
|
||||
for field, value in update_data.items():
|
||||
setattr(tool_type, field, value)
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(tool_type)
|
||||
return tool_type
|
||||
|
||||
|
||||
class ToolTypeValidateRequest(BaseModel):
|
||||
definition_type: str
|
||||
compose_template: str | None = None
|
||||
dockerfile_template: str | None = None
|
||||
|
||||
|
||||
@router.post(
|
||||
"/validate",
|
||||
summary="Validate tool type template",
|
||||
description="Validate a compose template or dockerfile syntax before creating a tool type.",
|
||||
)
|
||||
async def validate_tool_type_template(
|
||||
data: ToolTypeValidateRequest,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Validate a tool type template syntax.
|
||||
|
||||
Args:
|
||||
data: Validation request with definition type and template.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Validation result with success status and any errors.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
|
||||
errors = []
|
||||
|
||||
if data.definition_type == "compose":
|
||||
if not data.compose_template:
|
||||
errors.append("Compose template is required")
|
||||
else:
|
||||
try:
|
||||
validate_compose_yaml(data.compose_template)
|
||||
except ValueError as e:
|
||||
errors.append(str(e))
|
||||
|
||||
elif data.definition_type == "dockerfile":
|
||||
if not data.dockerfile_template:
|
||||
errors.append("Dockerfile template is required")
|
||||
elif not data.dockerfile_template.strip().startswith("FROM"):
|
||||
errors.append("Dockerfile must start with a FROM instruction")
|
||||
|
||||
else:
|
||||
errors.append("definition_type must be 'compose' or 'dockerfile'")
|
||||
|
||||
return {
|
||||
"valid": len(errors) == 0,
|
||||
"errors": errors,
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{tool_type_id}/validate",
|
||||
summary="Validate tool type",
|
||||
description="Validate the compose template or dockerfile syntax of a tool type.",
|
||||
)
|
||||
async def validate_tool_type(
|
||||
tool_type_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict:
|
||||
"""Validate a tool type's template syntax.
|
||||
|
||||
Args:
|
||||
tool_type_id: UUID of the tool type to validate.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
Validation result with success status and any errors.
|
||||
"""
|
||||
await _get_user(session, user_id)
|
||||
tool_type = await session.get(ToolType, tool_type_id)
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
|
||||
errors = []
|
||||
|
||||
if tool_type.definition_type == "compose":
|
||||
if not tool_type.compose_template:
|
||||
errors.append("Compose template is empty")
|
||||
else:
|
||||
try:
|
||||
validate_compose_yaml(tool_type.compose_template)
|
||||
except ValueError as e:
|
||||
errors.append(str(e))
|
||||
|
||||
elif tool_type.definition_type == "dockerfile":
|
||||
if not tool_type.dockerfile_template:
|
||||
errors.append("Dockerfile template is empty")
|
||||
elif not tool_type.dockerfile_template.strip().startswith("FROM"):
|
||||
errors.append("Dockerfile must start with a FROM instruction")
|
||||
|
||||
return {
|
||||
"valid": len(errors) == 0,
|
||||
"errors": errors,
|
||||
}
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/{tool_type_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="Delete tool type",
|
||||
description="Delete a custom tool type. Built-in tool types cannot be deleted.",
|
||||
)
|
||||
async def delete_tool_type(
|
||||
tool_type_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Delete a tool type.
|
||||
|
||||
Args:
|
||||
tool_type_id: UUID of the tool type to delete.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
None with 204 status code.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
await _require_admin(user)
|
||||
|
||||
tool_type = await session.get(ToolType, tool_type_id)
|
||||
if tool_type is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
||||
|
||||
# Built-in tool types can now be deleted
|
||||
|
||||
await session.delete(tool_type)
|
||||
await session.commit()
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Shared validation utilities for tool types."""
|
||||
|
||||
import re
|
||||
|
||||
import yaml
|
||||
from fastapi import HTTPException, status
|
||||
|
||||
|
||||
def sanitize_template_vars(template: str) -> str:
|
||||
"""Replace template variables like {{VAR}} with placeholders to avoid YAML parsing errors."""
|
||||
return re.sub(r"\{\{[A-Za-z_][A-Za-z0-9_]*\}\}", "__PLACEHOLDER__", template)
|
||||
|
||||
|
||||
def validate_compose_yaml(template: str) -> dict:
|
||||
"""Validate and parse a compose template.
|
||||
|
||||
Args:
|
||||
template: Raw compose template string.
|
||||
|
||||
Returns:
|
||||
Parsed YAML dict.
|
||||
|
||||
Raises:
|
||||
ValueError: If YAML is invalid or missing required keys.
|
||||
"""
|
||||
sanitized = sanitize_template_vars(template)
|
||||
|
||||
try:
|
||||
parsed = yaml.safe_load(sanitized)
|
||||
except yaml.YAMLError as e:
|
||||
raise ValueError(f"Invalid YAML: {e}")
|
||||
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("Compose template must be a YAML mapping")
|
||||
|
||||
if "services" not in parsed:
|
||||
raise ValueError("Compose template must contain 'services' key")
|
||||
|
||||
if not parsed["services"]:
|
||||
raise ValueError("Compose template must define at least one service")
|
||||
|
||||
return parsed
|
||||
|
||||
|
||||
def check_port_exposed(parsed: dict, port: int) -> bool:
|
||||
"""Check if a port is exposed in a parsed compose template.
|
||||
|
||||
Args:
|
||||
parsed: Parsed compose YAML dict.
|
||||
port: Port number to check.
|
||||
|
||||
Returns:
|
||||
True if port is exposed in any service.
|
||||
"""
|
||||
port_str = str(port)
|
||||
|
||||
if not isinstance(parsed, dict) or "services" not in parsed:
|
||||
return False
|
||||
|
||||
for service_config in parsed["services"].values():
|
||||
if isinstance(service_config, dict) and "ports" in service_config:
|
||||
for port_mapping in service_config["ports"]:
|
||||
if isinstance(port_mapping, str) and port_str in port_mapping:
|
||||
return True
|
||||
elif isinstance(port_mapping, int) and port_mapping == port:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def validate_required_variables(template: str, variables: list[str]) -> None:
|
||||
"""Validate that all required variables exist in the template.
|
||||
|
||||
Args:
|
||||
template: Compose template string.
|
||||
variables: List of required variable names.
|
||||
|
||||
Raises:
|
||||
HTTPException: If any variable is not found in the template.
|
||||
"""
|
||||
for var in variables:
|
||||
placeholder = f"{{{{{var}}}}}"
|
||||
if placeholder not in template:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Required variable '{var}' not found in compose template",
|
||||
)
|
||||
@@ -1,47 +1,29 @@
|
||||
import logging
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from typing import Any
|
||||
|
||||
from src.auth.jwt_service import decode_access_token
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.user import User
|
||||
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/users/me", tags=["user-config"])
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
access_token: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
try:
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return uuid.UUID(str(claims["sub"]))
|
||||
except Exception:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid access token")
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
async def _get_or_create_config(session: AsyncSession, user_id: uuid.UUID) -> UserConfig:
|
||||
"""Get or create user config record.
|
||||
|
||||
Args:
|
||||
session: Database session.
|
||||
user_id: UUID of the user.
|
||||
|
||||
Returns:
|
||||
The user's config, creating a new one if it doesn't exist.
|
||||
"""
|
||||
result = await session.execute(select(UserConfig).where(UserConfig.user_id == user_id))
|
||||
config = result.scalar_one_or_none()
|
||||
if config is None:
|
||||
@@ -59,6 +41,7 @@ class UserConfigResponse(BaseModel):
|
||||
theme: str = "system"
|
||||
git_user_name: str | None = None
|
||||
git_user_email: str | None = None
|
||||
last_session_id: str | None = None
|
||||
|
||||
|
||||
class UserConfigUpdate(BaseModel):
|
||||
@@ -66,31 +49,64 @@ class UserConfigUpdate(BaseModel):
|
||||
theme: str | None = None
|
||||
git_user_name: str | None = None
|
||||
git_user_email: str | None = None
|
||||
last_session_id: str | None = None
|
||||
|
||||
|
||||
@router.get("/config", response_model=UserConfigResponse)
|
||||
@router.get(
|
||||
"/config",
|
||||
response_model=UserConfigResponse,
|
||||
summary="Get user config",
|
||||
description="Get the current user's configuration settings.",
|
||||
)
|
||||
async def get_user_config(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> UserConfigResponse:
|
||||
"""Get the current user's configuration.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The user's configuration settings.
|
||||
"""
|
||||
_user = await _get_user(session, user_id)
|
||||
config = await _get_or_create_config(session, user_id)
|
||||
return UserConfigResponse.model_validate(config.config)
|
||||
|
||||
|
||||
@router.patch("/config", response_model=UserConfigResponse)
|
||||
@router.patch(
|
||||
"/config",
|
||||
response_model=UserConfigResponse,
|
||||
summary="Update user config",
|
||||
description="Update the current user's configuration settings.",
|
||||
)
|
||||
async def update_user_config(
|
||||
data: UserConfigUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> UserConfigResponse:
|
||||
"""Update the current user's configuration.
|
||||
|
||||
Args:
|
||||
data: Configuration update data with optional fields.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated user configuration.
|
||||
"""
|
||||
_user = await _get_user(session, user_id)
|
||||
config = await _get_or_create_config(session, user_id)
|
||||
|
||||
# Merge updates
|
||||
update_data = data.model_dump(exclude_unset=True, exclude_none=True)
|
||||
config.config.update(update_data)
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
logger.debug("Updating user config for user %s: %s", user_id, update_data)
|
||||
# SQLAlchemy JSON doesn't track dict mutations, so we replace the whole dict
|
||||
config.config = {**config.config, **update_data}
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(config)
|
||||
logger.debug("Updated config: %s", config.config)
|
||||
return UserConfigResponse.model_validate(config.config)
|
||||
|
||||
+49
-33
@@ -1,14 +1,11 @@
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, UploadFile, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, UploadFile, status
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.jwt_service import decode_access_token
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/users", tags=["users"])
|
||||
@@ -19,31 +16,6 @@ ALLOWED_CONTENT_TYPES = {"image/png", "image/jpeg", "image/jpg"}
|
||||
MAX_AVATAR_SIZE = 2 * 1024 * 1024 # 2MB
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
access_token: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
try:
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return uuid.UUID(str(claims["sub"]))
|
||||
except Exception:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid access token")
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
class UserProfileResponse(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -58,20 +30,49 @@ class UserProfileUpdate(BaseModel):
|
||||
email: str | None = None
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserProfileResponse)
|
||||
@router.get(
|
||||
"/me",
|
||||
response_model=UserProfileResponse,
|
||||
summary="Get current user profile",
|
||||
description="Retrieve the profile of the currently authenticated user.",
|
||||
)
|
||||
async def get_profile(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
"""Get the current user's profile.
|
||||
|
||||
Args:
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The user's profile information.
|
||||
"""
|
||||
return await _get_user(session, user_id)
|
||||
|
||||
|
||||
@router.put("/me", response_model=UserProfileResponse)
|
||||
@router.put(
|
||||
"/me",
|
||||
response_model=UserProfileResponse,
|
||||
summary="Update user profile",
|
||||
description="Update the current user's profile information.",
|
||||
)
|
||||
async def update_profile(
|
||||
data: UserProfileUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
"""Update the current user's profile.
|
||||
|
||||
Args:
|
||||
data: Profile update data with optional name and email.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated user profile.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
|
||||
if data.name is not None:
|
||||
@@ -89,12 +90,27 @@ async def update_profile(
|
||||
return user
|
||||
|
||||
|
||||
@router.post("/me/avatar", response_model=UserProfileResponse)
|
||||
@router.post(
|
||||
"/me/avatar",
|
||||
response_model=UserProfileResponse,
|
||||
summary="Upload avatar",
|
||||
description="Upload a profile avatar image (PNG or JPG, max 2MB).",
|
||||
)
|
||||
async def upload_avatar(
|
||||
file: UploadFile,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
"""Upload a profile avatar image.
|
||||
|
||||
Args:
|
||||
file: The image file to upload (PNG or JPG, max 2MB).
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The updated user profile with new avatar URL.
|
||||
"""
|
||||
user = await _get_user(session, user_id)
|
||||
|
||||
if file.content_type not in ALLOWED_CONTENT_TYPES:
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import build_login_redirect_url
|
||||
from src.auth.refresh_store import hash_refresh_token
|
||||
from src.auth.session import create_session_cookie, decode_session_cookie
|
||||
|
||||
__all__ = [
|
||||
"build_cookie_options",
|
||||
"build_login_redirect_url",
|
||||
"decode_access_token",
|
||||
"hash_refresh_token",
|
||||
"mint_access_token",
|
||||
"create_session_cookie",
|
||||
"decode_session_cookie",
|
||||
]
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def build_cookie_options(settings: Settings) -> dict[str, str | bool]:
|
||||
def build_cookie_options(settings: Settings) -> dict[str, str | bool | None]:
|
||||
return {
|
||||
"httponly": True,
|
||||
"secure": settings.cookie_secure,
|
||||
"samesite": settings.cookie_samesite,
|
||||
"domain": settings.cookie_domain,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import Cookie, Depends, HTTPException, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.session import decode_session_cookie
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.project import Project
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
session_cookie: Annotated[str | None, Cookie(alias="session")] = None,
|
||||
) -> uuid.UUID:
|
||||
if not session_cookie:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
return uuid.UUID(str(payload["user_id"]))
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid session")
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
session_cookie: Annotated[str | None, Cookie(alias="session")] = None,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
if not session_cookie:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
user_id = uuid.UUID(str(payload["user_id"]))
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid session")
|
||||
|
||||
user = await db_session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
"""Fetch a user by ID or raise 401 if not found."""
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
async def _get_owned_project(
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
) -> "Project":
|
||||
"""Fetch a project and verify ownership.
|
||||
|
||||
Args:
|
||||
project_id: UUID of the project.
|
||||
user_id: ID of the authenticated user.
|
||||
session: Database session.
|
||||
|
||||
Returns:
|
||||
The project if found and owned by the user.
|
||||
|
||||
Raises:
|
||||
HTTPException: 404 if project not found, 403 if user is not the owner.
|
||||
"""
|
||||
from src.models.project import Project
|
||||
|
||||
project = await session.get(Project, project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="project not found")
|
||||
if project.owner_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="not project owner")
|
||||
return project
|
||||
@@ -1,27 +0,0 @@
|
||||
from datetime import datetime
|
||||
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def mint_access_token(
|
||||
*,
|
||||
settings: Settings,
|
||||
subject: str,
|
||||
email: str,
|
||||
name: str,
|
||||
expires_at: datetime,
|
||||
) -> str:
|
||||
payload = {
|
||||
"sub": subject,
|
||||
"email": email,
|
||||
"name": name,
|
||||
"exp": expires_at,
|
||||
}
|
||||
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
|
||||
|
||||
|
||||
def decode_access_token(*, settings: Settings, token: str) -> dict[str, str | int]:
|
||||
claims = jwt.decode(token, settings.jwt_secret, algorithms=[settings.jwt_algorithm])
|
||||
return dict(claims)
|
||||
+12
-25
@@ -1,7 +1,7 @@
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
@@ -11,7 +11,6 @@ def build_login_redirect_url(
|
||||
settings: Settings,
|
||||
redirect_uri: str,
|
||||
state: str,
|
||||
nonce: str,
|
||||
) -> str:
|
||||
query = urlencode(
|
||||
{
|
||||
@@ -20,7 +19,6 @@ def build_login_redirect_url(
|
||||
"redirect_uri": redirect_uri,
|
||||
"scope": "openid profile email",
|
||||
"state": state,
|
||||
"nonce": nonce,
|
||||
}
|
||||
)
|
||||
return f"{settings.resolved_authentik_authorize_url}?{query}"
|
||||
@@ -47,31 +45,20 @@ async def exchange_code_for_tokens(
|
||||
payload = response.json()
|
||||
return {
|
||||
"access_token": payload["access_token"],
|
||||
"refresh_token": payload["refresh_token"],
|
||||
"refresh_token": payload.get("refresh_token"),
|
||||
}
|
||||
|
||||
|
||||
async def fetch_jwks(*, settings: Settings, client: httpx.AsyncClient) -> dict[str, list[dict[str, str]]]:
|
||||
response = await client.get(settings.resolved_authentik_jwks_url)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
return {"keys": payload["keys"]}
|
||||
|
||||
|
||||
def verify_provider_access_token(
|
||||
async def fetch_user_info(
|
||||
*,
|
||||
settings: Settings,
|
||||
token: str,
|
||||
jwks: dict[str, list[dict[str, str]]],
|
||||
) -> dict[str, str | int]:
|
||||
unverified_header = jwt.get_unverified_header(token)
|
||||
key_id = unverified_header["kid"]
|
||||
jwk_key = next(key for key in jwks["keys"] if key.get("kid") == key_id)
|
||||
claims = jwt.decode(
|
||||
token,
|
||||
jwk_key,
|
||||
algorithms=[jwk_key.get("alg", "HS256")],
|
||||
audience=settings.authentik_audience,
|
||||
issuer=settings.resolved_authentik_issuer,
|
||||
access_token: str,
|
||||
client: httpx.AsyncClient,
|
||||
) -> dict[str, Any]:
|
||||
"""Fetch user info from Authentik userinfo endpoint."""
|
||||
response = await client.get(
|
||||
f"{settings.authentik_base_url}/application/o/userinfo/",
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
)
|
||||
return dict(claims)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
from datetime import UTC, datetime
|
||||
from hashlib import sha256
|
||||
from secrets import token_urlsafe
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.models.refresh_token import RefreshToken
|
||||
|
||||
|
||||
def hash_refresh_token(raw_token: str) -> str:
|
||||
return sha256(raw_token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def create_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
user_id: object,
|
||||
expires_at: datetime,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
raw_token = token_urlsafe(48)
|
||||
record = RefreshToken(
|
||||
user_id=user_id,
|
||||
token_hash=hash_refresh_token(raw_token),
|
||||
expires_at=expires_at,
|
||||
created_at=datetime.now(UTC),
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
await session.refresh(record)
|
||||
return raw_token, record
|
||||
|
||||
|
||||
async def rotate_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
raw_token: str,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
existing_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(
|
||||
select(RefreshToken).where(
|
||||
RefreshToken.token_hash == existing_hash,
|
||||
RefreshToken.revoked_at.is_(None),
|
||||
)
|
||||
)
|
||||
if existing is None:
|
||||
raise ValueError("refresh token not found")
|
||||
if existing.expires_at <= datetime.now(UTC):
|
||||
raise ValueError("refresh token expired")
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.flush()
|
||||
|
||||
return await create_refresh_token(
|
||||
session=session,
|
||||
user_id=existing.user_id,
|
||||
expires_at=existing.expires_at,
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
|
||||
|
||||
async def revoke_refresh_token(*, session: AsyncSession, raw_token: str) -> bool:
|
||||
token_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
if existing is None:
|
||||
return False
|
||||
if existing.revoked_at is not None:
|
||||
return True
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.commit()
|
||||
return True
|
||||
@@ -0,0 +1,71 @@
|
||||
import hmac
|
||||
import hashlib
|
||||
import json
|
||||
import base64
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def _base64url_encode(data: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _base64url_decode(data: str) -> bytes:
|
||||
padding = 4 - len(data) % 4
|
||||
if padding != 4:
|
||||
data += "=" * padding
|
||||
return base64.urlsafe_b64decode(data)
|
||||
|
||||
|
||||
def create_session_cookie(*, settings: Settings, user_id: str) -> str:
|
||||
"""Create a signed session cookie value."""
|
||||
payload = {
|
||||
"user_id": user_id,
|
||||
"exp": int((datetime.now(timezone.utc) + timedelta(hours=settings.session_ttl_hours)).timestamp()),
|
||||
}
|
||||
|
||||
header = _base64url_encode(json.dumps({"alg": "HS256", "typ": "session"}).encode())
|
||||
payload_encoded = _base64url_encode(json.dumps(payload).encode())
|
||||
message = f"{header}.{payload_encoded}"
|
||||
|
||||
signature = hmac.new(
|
||||
settings.session_secret.encode(),
|
||||
message.encode(),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
signature_encoded = _base64url_encode(signature)
|
||||
|
||||
return f"{message}.{signature_encoded}"
|
||||
|
||||
|
||||
def decode_session_cookie(*, settings: Settings, cookie_value: str) -> dict[str, Any]:
|
||||
"""Decode and verify a session cookie. Returns payload or raises ValueError."""
|
||||
parts = cookie_value.split(".")
|
||||
if len(parts) != 3:
|
||||
raise ValueError("invalid session format")
|
||||
|
||||
header, payload_encoded, signature_encoded = parts
|
||||
message = f"{header}.{payload_encoded}"
|
||||
|
||||
# Verify signature
|
||||
expected_signature = hmac.new(
|
||||
settings.session_secret.encode(),
|
||||
message.encode(),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
expected_signature_encoded = _base64url_encode(expected_signature)
|
||||
|
||||
if not hmac.compare_digest(signature_encoded, expected_signature_encoded):
|
||||
raise ValueError("invalid session signature")
|
||||
|
||||
# Decode payload
|
||||
payload_bytes = _base64url_decode(payload_encoded)
|
||||
payload = json.loads(payload_bytes)
|
||||
|
||||
# Check expiry
|
||||
if payload.get("exp", 0) < int(datetime.now(timezone.utc).timestamp()):
|
||||
raise ValueError("session expired")
|
||||
|
||||
return payload
|
||||
+30
-7
@@ -34,19 +34,26 @@ class Settings(BaseSettings):
|
||||
# Authentik configuration - no hardcoded URLs
|
||||
authentik_client_id: str = "headquarter-web"
|
||||
authentik_client_secret: str = "change-me"
|
||||
# Authentik application slug used in URLs (e.g., "headquarter-web")
|
||||
# This is different from the OAuth client_id which may be a UUID
|
||||
authentik_application_slug: str = "headquarter-web"
|
||||
authentik_authorize_url: str | None = None
|
||||
authentik_token_url: str | None = None
|
||||
authentik_jwks_url: str | None = None
|
||||
authentik_issuer: str | None = None
|
||||
authentik_audience: str = "headquarter-web"
|
||||
|
||||
jwt_secret: str = "change-me-jwt-secret"
|
||||
jwt_algorithm: str = "HS256"
|
||||
access_token_ttl_minutes: int = 15
|
||||
refresh_token_ttl_days: int = 7
|
||||
# Session configuration
|
||||
session_secret: str = "change-me-session-secret"
|
||||
session_ttl_hours: int = 24
|
||||
|
||||
# Repository storage
|
||||
repo_base_path: str = "/data/repos"
|
||||
|
||||
# Tool instance storage
|
||||
instance_base_path: str = "/data/instances"
|
||||
|
||||
|
||||
|
||||
model_config = SettingsConfigDict(env_file=".env", extra="ignore", populate_by_name=True)
|
||||
|
||||
@@ -100,13 +107,13 @@ class Settings(BaseSettings):
|
||||
def resolved_authentik_jwks_url(self) -> str:
|
||||
if self.authentik_jwks_url:
|
||||
return self.authentik_jwks_url
|
||||
return f"{self.authentik_base_url}/application/o/{self.authentik_client_id}/jwks/"
|
||||
return f"{self.authentik_base_url}/application/o/{self.authentik_application_slug}/jwks/"
|
||||
|
||||
@property
|
||||
def resolved_authentik_issuer(self) -> str:
|
||||
if self.authentik_issuer:
|
||||
return self.authentik_issuer
|
||||
return f"{self.authentik_base_url}/application/o/{self.authentik_client_id}/"
|
||||
return f"{self.authentik_base_url}/application/o/{self.authentik_application_slug}/"
|
||||
|
||||
@property
|
||||
def cookie_secure(self) -> bool:
|
||||
@@ -115,6 +122,22 @@ class Settings(BaseSettings):
|
||||
@property
|
||||
def cookie_samesite(self) -> str:
|
||||
if self.app_env == "production":
|
||||
return "strict"
|
||||
return "none"
|
||||
|
||||
return "lax"
|
||||
|
||||
@property
|
||||
def cookie_domain(self) -> str | None:
|
||||
"""Return the parent domain for cross-subdomain cookies.
|
||||
|
||||
E.g., api.example.com and app.example.com both share .example.com
|
||||
"""
|
||||
if self.app_env != "production":
|
||||
return None
|
||||
|
||||
# Extract parent domain from api_domain
|
||||
# e.g., "api.headquarter.commumedia.org" -> ".headquarter.commumedia.org"
|
||||
parts = self.api_domain.split(".")
|
||||
if len(parts) >= 3:
|
||||
return "." + ".".join(parts[1:])
|
||||
return None
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import subprocess
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from src.config import Settings, build_database_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
settings = Settings()
|
||||
database_url = settings.database_url
|
||||
@@ -20,4 +25,92 @@ engine = create_async_engine(
|
||||
)
|
||||
SessionLocal = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
|
||||
|
||||
__all__ = ["SessionLocal", "build_database_url", "engine", "settings"]
|
||||
|
||||
async def init_database(
|
||||
max_retries: int = 5,
|
||||
retry_delay: float = 2.0,
|
||||
) -> bool:
|
||||
"""Initialize the database by running pending migrations.
|
||||
|
||||
Uses subprocess to run 'alembic upgrade head' to avoid
|
||||
async/sync context manager issues with SQLAlchemy 2.0.
|
||||
|
||||
Returns True if migrations succeeded, False otherwise.
|
||||
"""
|
||||
for attempt in range(1, max_retries + 1):
|
||||
try:
|
||||
# Test basic connectivity
|
||||
from sqlalchemy import text
|
||||
test_conn = await engine.connect()
|
||||
try:
|
||||
await test_conn.execute(text("SELECT 1"))
|
||||
finally:
|
||||
await test_conn.close()
|
||||
|
||||
logger.info("Database connection established.")
|
||||
|
||||
# Run migrations via subprocess
|
||||
logger.info("Running database migrations...")
|
||||
result = await asyncio.get_event_loop().run_in_executor(
|
||||
None,
|
||||
lambda: subprocess.run(
|
||||
["alembic", "upgrade", "head"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd="/app",
|
||||
),
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
logger.info("Database migrations completed successfully.")
|
||||
logger.debug("Alembic output: %s", result.stdout)
|
||||
return True
|
||||
else:
|
||||
logger.error("Migration failed: %s", result.stderr)
|
||||
if attempt < max_retries:
|
||||
wait = retry_delay * (2 ** (attempt - 1))
|
||||
logger.info("Retrying in %.1f seconds...", wait)
|
||||
await asyncio.sleep(wait)
|
||||
else:
|
||||
return False
|
||||
|
||||
except Exception as exc:
|
||||
error_msg = str(exc).lower()
|
||||
if "connection" in error_msg or "could not connect" in error_msg:
|
||||
logger.warning(
|
||||
"Database connection failed (attempt %d/%d): %s",
|
||||
attempt,
|
||||
max_retries,
|
||||
exc,
|
||||
)
|
||||
elif "authentication" in error_msg or "password" in error_msg:
|
||||
logger.error(
|
||||
"Database authentication failed: %s. "
|
||||
"Check POSTGRES_USER and POSTGRES_PASSWORD environment variables.",
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
else:
|
||||
logger.error(
|
||||
"Database initialization error (attempt %d/%d): %s",
|
||||
attempt,
|
||||
max_retries,
|
||||
exc,
|
||||
)
|
||||
|
||||
if attempt < max_retries:
|
||||
wait = retry_delay * (2 ** (attempt - 1))
|
||||
logger.info("Retrying in %.1f seconds...", wait)
|
||||
await asyncio.sleep(wait)
|
||||
else:
|
||||
logger.error(
|
||||
"Failed to initialize database after %d attempts. "
|
||||
"Ensure the database is running and accessible.",
|
||||
max_retries,
|
||||
)
|
||||
return False
|
||||
|
||||
return False
|
||||
|
||||
|
||||
__all__ = ["SessionLocal", "build_database_url", "engine", "settings", "init_database"]
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import logging
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from typing import Callable
|
||||
|
||||
from fastapi import Request, Response
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""Log all HTTP requests with timing and status codes."""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
start_time = time.time()
|
||||
client_host = request.client.host if request.client else "unknown"
|
||||
|
||||
# Log the incoming request
|
||||
logger.info(
|
||||
"→ Request: %s %s (client: %s)",
|
||||
request.method,
|
||||
request.url.path,
|
||||
client_host,
|
||||
)
|
||||
|
||||
try:
|
||||
response = await call_next(request)
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Log the response
|
||||
logger.info(
|
||||
"← Response: %s %s → %d (%dms)",
|
||||
request.method,
|
||||
request.url.path,
|
||||
response.status_code,
|
||||
int(duration * 1000),
|
||||
)
|
||||
return response
|
||||
|
||||
except Exception as exc:
|
||||
duration = time.time() - start_time
|
||||
logger.error(
|
||||
"✗ Error: %s %s → %s (%dms)\n%s",
|
||||
request.method,
|
||||
request.url.path,
|
||||
type(exc).__name__,
|
||||
int(duration * 1000),
|
||||
traceback.format_exc(),
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
class ExceptionLoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""Catch and log all unhandled exceptions."""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
try:
|
||||
return await call_next(request)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Unhandled exception in %s %s:\n%s",
|
||||
request.method,
|
||||
request.url.path,
|
||||
traceback.format_exc(),
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def configure_logging(level: int = logging.INFO) -> None:
|
||||
"""Configure structured logging for the application."""
|
||||
formatter = logging.Formatter(
|
||||
fmt="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
|
||||
# Console handler
|
||||
console_handler = logging.StreamHandler(sys.stdout)
|
||||
console_handler.setFormatter(formatter)
|
||||
|
||||
# Configure root logger
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.setLevel(level)
|
||||
root_logger.handlers = [console_handler]
|
||||
|
||||
# Set levels for specific loggers
|
||||
logging.getLogger("uvicorn").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("sqlalchemy.engine").setLevel(logging.WARNING)
|
||||
|
||||
logger.info("Logging configured at level %s", logging.getLevelName(level))
|
||||
+116
-1
@@ -1,18 +1,133 @@
|
||||
from fastapi import FastAPI
|
||||
import logging
|
||||
import os
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from src.api.auth import router as auth_router
|
||||
from src.api.dashboard import router as dashboard_router
|
||||
from src.api.git_repositories import router as git_repositories_router
|
||||
from src.api.health import router as health_router
|
||||
from src.api.projects import router as projects_router
|
||||
from src.api.ssh_keys import router as ssh_keys_router
|
||||
from src.api.terminal import router as terminal_router
|
||||
from src.api.instance_proxy import router as instance_proxy_router
|
||||
from src.api.config_folders import router as config_folders_router
|
||||
from src.api.config_profiles import router as config_profiles_router
|
||||
from src.api.tool_configs import router as tool_configs_router
|
||||
from src.api.tool_instances import router as tool_instances_router
|
||||
from src.api.tool_instances import sessions_router
|
||||
from src.api.tool_types import router as tool_types_router
|
||||
from src.api.user_config import router as user_config_router
|
||||
from src.api.users import router as users_router
|
||||
from src.config import Settings
|
||||
from src.database import init_database
|
||||
from src.logging_config import (
|
||||
ExceptionLoggingMiddleware,
|
||||
RequestLoggingMiddleware,
|
||||
configure_logging,
|
||||
)
|
||||
|
||||
# Configure logging early
|
||||
log_level = os.getenv("LOG_LEVEL", "INFO").upper()
|
||||
configure_logging(level=getattr(logging, log_level, logging.INFO))
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = Settings()
|
||||
app = FastAPI(title="Headquarter API")
|
||||
|
||||
# Configure CORS - must be before other middleware
|
||||
# Build allowed origins list including web and api domains
|
||||
cors_origins = [settings.web_base_url]
|
||||
if settings.api_base_url != settings.web_base_url:
|
||||
cors_origins.append(settings.api_base_url)
|
||||
logger.info("CORS configured with origins: %s", cors_origins)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=cors_origins,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.add_middleware(RequestLoggingMiddleware)
|
||||
app.add_middleware(ExceptionLoggingMiddleware)
|
||||
|
||||
|
||||
def _sanitize_validation_errors(errors):
|
||||
"""Convert validation errors to JSON-safe format."""
|
||||
sanitized = []
|
||||
for error in errors:
|
||||
safe_error = {
|
||||
"type": error.get("type"),
|
||||
"loc": error.get("loc"),
|
||||
"msg": error.get("msg"),
|
||||
"input": str(error.get("input")) if error.get("input") is not None else None,
|
||||
}
|
||||
# Convert ctx to safe format
|
||||
ctx = error.get("ctx")
|
||||
if ctx:
|
||||
safe_ctx = {}
|
||||
for key, value in ctx.items():
|
||||
if isinstance(value, Exception):
|
||||
safe_ctx[key] = str(value)
|
||||
elif isinstance(value, (str, int, float, bool, type(None))):
|
||||
safe_ctx[key] = value
|
||||
else:
|
||||
safe_ctx[key] = str(value)
|
||||
safe_error["ctx"] = safe_ctx
|
||||
sanitized.append(safe_error)
|
||||
return sanitized
|
||||
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_exception_handler(request: Request, exc: RequestValidationError):
|
||||
"""Log validation errors and return detailed response."""
|
||||
errors = exc.errors()
|
||||
logger.warning(
|
||||
"Validation error for %s %s: %s",
|
||||
request.method,
|
||||
request.url.path,
|
||||
errors,
|
||||
)
|
||||
safe_errors = _sanitize_validation_errors(errors)
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
content={"detail": safe_errors},
|
||||
)
|
||||
|
||||
|
||||
@app.on_event("startup")
|
||||
async def on_startup():
|
||||
logger.info("Starting up Headquarter API...")
|
||||
|
||||
# Initialize database (run migrations)
|
||||
db_ready = await init_database()
|
||||
if not db_ready:
|
||||
logger.error("Database initialization failed. Shutting down.")
|
||||
import sys
|
||||
sys.exit(1)
|
||||
|
||||
logger.info("Startup complete.")
|
||||
|
||||
app.include_router(health_router)
|
||||
app.include_router(auth_router)
|
||||
app.include_router(dashboard_router)
|
||||
app.include_router(projects_router)
|
||||
app.include_router(users_router)
|
||||
app.include_router(ssh_keys_router)
|
||||
app.include_router(git_repositories_router)
|
||||
app.include_router(user_config_router)
|
||||
app.include_router(tool_types_router)
|
||||
app.include_router(config_folders_router)
|
||||
app.include_router(config_profiles_router)
|
||||
app.include_router(tool_instances_router)
|
||||
app.include_router(tool_configs_router)
|
||||
app.include_router(sessions_router)
|
||||
app.include_router(instance_proxy_router)
|
||||
app.include_router(terminal_router)
|
||||
app.mount("/uploads", StaticFiles(directory="uploads"), name="uploads")
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
from src.models.base import Base
|
||||
from src.models.config_folder import ConfigFolder
|
||||
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.tool_instance import ToolInstance
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
__all__ = ["Base", "GitRepository", "Project", "RefreshToken", "SSHKey", "User", "UserConfig"]
|
||||
__all__ = ["Base", "ConfigFolder", "ConfigProfile", "ConfigProfileInclude", "GitRepository", "Project", "SSHKey", "ToolInstance", "ToolType", "User", "UserConfig"]
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, JSON, String, Text
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class ConfigFolder(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "config_folders"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("users.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
mount_path: Mapped[str] = mapped_column(String(1024), nullable=False)
|
||||
files: Mapped[dict] = mapped_column(
|
||||
JSON, default=dict, nullable=False
|
||||
) # {"relative/path": "content", ...}
|
||||
project_overrides: Mapped[dict | None] = mapped_column(
|
||||
JSON, default=dict, nullable=True
|
||||
) # {"project_id": {"mount_path": "...", "files": {...}}}
|
||||
is_active: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
||||
|
||||
user: Mapped["User"] = relationship()
|
||||
@@ -0,0 +1,77 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import ForeignKey, JSON, Integer, String, Text, Boolean
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class ConfigProfile(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "config_profiles"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("users.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
project_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(), ForeignKey("projects.id", ondelete="CASCADE"), nullable=True
|
||||
)
|
||||
tool_type_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(), ForeignKey("tool_types.id", ondelete="CASCADE"), nullable=True
|
||||
)
|
||||
env_vars: Mapped[dict] = mapped_column(
|
||||
JSON, default=dict, nullable=False
|
||||
) # {"VAR_NAME": "value", ...}
|
||||
runtime_hints: Mapped[dict] = mapped_column(
|
||||
JSON, default=dict, nullable=False
|
||||
) # {"start_command": "...", "working_dir": "...", ...}
|
||||
mounts: Mapped[list] = mapped_column(
|
||||
JSON, default=list, nullable=False
|
||||
) # [{"target": "/path", "mode": "rw", "files": {"rel/path": "content"}}, ...]
|
||||
files: Mapped[dict] = mapped_column(
|
||||
JSON, default=dict, nullable=False
|
||||
) # {"rel/path": "content", ...}
|
||||
git_mounts: Mapped[list] = mapped_column(
|
||||
JSON, default=list, nullable=False
|
||||
) # [{"remote_url": "https://github.com/user/repo.git", "source_path": ".", "target_path": "/path", "branch": "main"}, ...]
|
||||
is_default: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
|
||||
user: Mapped["User"] = relationship()
|
||||
project: Mapped["Project | None"] = relationship()
|
||||
tool_type: Mapped["ToolType | None"] = relationship()
|
||||
includes: Mapped[list["ConfigProfileInclude"]] = relationship(
|
||||
"ConfigProfileInclude",
|
||||
foreign_keys="ConfigProfileInclude.profile_id",
|
||||
order_by="ConfigProfileInclude.order_index",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
|
||||
class ConfigProfileInclude(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "config_profile_includes"
|
||||
|
||||
profile_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("config_profiles.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
included_profile_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("config_profiles.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
order_index: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
|
||||
profile: Mapped["ConfigProfile"] = relationship(
|
||||
"ConfigProfile",
|
||||
foreign_keys=[profile_id],
|
||||
back_populates="includes",
|
||||
)
|
||||
included_profile: Mapped["ConfigProfile"] = relationship(
|
||||
"ConfigProfile",
|
||||
foreign_keys=[included_profile_id],
|
||||
)
|
||||
@@ -10,6 +10,7 @@ from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
@@ -18,11 +19,15 @@ class GitRepository(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
path: Mapped[str] = mapped_column(String(1024))
|
||||
project_id: Mapped[uuid.UUID] = mapped_column(UUID(), ForeignKey("projects.id"), nullable=False)
|
||||
project_id: Mapped[uuid.UUID | None] = mapped_column(UUID(), ForeignKey("projects.id"), nullable=True)
|
||||
owner_id: Mapped[uuid.UUID] = mapped_column(UUID(), ForeignKey("users.id"), nullable=False)
|
||||
is_mirror: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
remote_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
last_push: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
ssh_key_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(), ForeignKey("ssh_keys.id"), nullable=True
|
||||
)
|
||||
|
||||
project: Mapped["Project"] = relationship(back_populates="repositories")
|
||||
owner: Mapped["User"] = relationship()
|
||||
ssh_key: Mapped["SSHKey | None"] = relationship()
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class RefreshToken(UUIDPrimaryKeyMixin, Base):
|
||||
__tablename__ = "refresh_tokens"
|
||||
|
||||
user_id: Mapped[str] = mapped_column(ForeignKey("users.id"), nullable=False, index=True)
|
||||
token_hash: Mapped[str] = mapped_column(String(255), unique=True, nullable=False)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, index=True)
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
user_agent: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
ip_address: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
user: Mapped["User"] = relationship(back_populates="refresh_tokens")
|
||||
@@ -5,14 +5,14 @@ from sqlalchemy import ForeignKey, String, Text
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class SSHKey(UUIDPrimaryKeyMixin, Base):
|
||||
class SSHKey(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "ssh_keys"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import ForeignKey, JSON, String, Text
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class ToolConfig(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "tool_configs"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("users.id"), nullable=False
|
||||
)
|
||||
tool_type_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("tool_types.id"), nullable=False
|
||||
)
|
||||
project_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(), ForeignKey("projects.id"), nullable=True
|
||||
)
|
||||
key: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
value: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
config_type: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False, default="env"
|
||||
) # "env" or "file"
|
||||
file_path: Mapped[str | None] = mapped_column(
|
||||
String(1024), nullable=True
|
||||
) # Only for file type
|
||||
port_override: Mapped[int | None] = mapped_column(nullable=True)
|
||||
start_command: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
working_directory: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
environment_variables: Mapped[dict | None] = mapped_column(
|
||||
JSON, default=dict, nullable=True
|
||||
)
|
||||
volumes: Mapped[list[dict] | None] = mapped_column(
|
||||
JSON, default=list, nullable=True
|
||||
)
|
||||
|
||||
user: Mapped["User"] = relationship()
|
||||
tool_type: Mapped["ToolType"] = relationship()
|
||||
project: Mapped["Project | None"] = relationship()
|
||||
@@ -0,0 +1,83 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, Integer, JSON, String
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.config_profile import ConfigProfile
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class ToolInstance(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "tool_instances"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
display_name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
tool_type_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("tool_types.id"), nullable=False
|
||||
)
|
||||
repository_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("git_repositories.id"), nullable=False
|
||||
)
|
||||
project_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("projects.id"), nullable=False
|
||||
)
|
||||
owner_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(), ForeignKey("users.id"), nullable=False
|
||||
)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(50), nullable=False, default="pending"
|
||||
)
|
||||
container_id: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
container_name: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
compose_path: Mapped[str | None] = mapped_column(
|
||||
String(1024), nullable=True
|
||||
)
|
||||
url: Mapped[str | None] = mapped_column(
|
||||
String(1024), nullable=True
|
||||
)
|
||||
public_url: Mapped[str | None] = mapped_column(
|
||||
String(1024), nullable=True
|
||||
)
|
||||
tunnel_id: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
port: Mapped[int | None] = mapped_column(
|
||||
Integer, nullable=True
|
||||
)
|
||||
last_started_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
last_stopped_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
probe_result: Mapped[dict | None] = mapped_column(
|
||||
JSON, nullable=True
|
||||
)
|
||||
clone_mode: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False, default="mount"
|
||||
)
|
||||
branch: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True, default="main"
|
||||
)
|
||||
selected_config_profile_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(), ForeignKey("config_profiles.id", ondelete="SET NULL"), nullable=True
|
||||
)
|
||||
|
||||
tool_type: Mapped["ToolType"] = relationship()
|
||||
repository: Mapped["GitRepository"] = relationship()
|
||||
project: Mapped["Project"] = relationship()
|
||||
owner: Mapped["User"] = relationship()
|
||||
selected_config_profile: Mapped["ConfigProfile | None"] = relationship()
|
||||
@@ -0,0 +1,41 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, JSON, String, Text
|
||||
from sqlalchemy import Uuid as UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class ToolType(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "tool_types"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255), unique=True, nullable=False)
|
||||
display_name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
category: Mapped[str] = mapped_column(String(50), nullable=False, default="other")
|
||||
interface_type: Mapped[str] = mapped_column(String(20), nullable=False, default="web")
|
||||
requires_port: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
||||
default_port: Mapped[int] = mapped_column(nullable=False)
|
||||
definition_type: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False, default="compose"
|
||||
) # "compose" or "dockerfile"
|
||||
compose_template: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
dockerfile_template: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
build_context: Mapped[dict | None] = mapped_column(
|
||||
JSON, default=dict, nullable=True
|
||||
)
|
||||
readiness_probe: Mapped[dict | None] = mapped_column(JSON, nullable=True)
|
||||
startup_command: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
required_variables: Mapped[list[str]] = mapped_column(JSON, default=list, nullable=False)
|
||||
created_by_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(),
|
||||
ForeignKey("users.id"),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
created_by: Mapped["User | None"] = relationship()
|
||||
@@ -7,7 +7,6 @@ from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
@@ -21,6 +20,5 @@ class User(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
avatar_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
|
||||
projects: Mapped[list["Project"]] = relationship(back_populates="owner")
|
||||
refresh_tokens: Mapped[list["RefreshToken"]] = relationship(back_populates="user")
|
||||
ssh_keys: Mapped[list["SSHKey"]] = relationship(back_populates="user")
|
||||
user_config: Mapped["UserConfig | None"] = relationship(back_populates="user", uselist=False)
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Clone service for repository cloning and dirty state checking."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def clone_repository(
|
||||
remote_url: str,
|
||||
ssh_key_path: str | None,
|
||||
instance_dir: str,
|
||||
branch: str = "main",
|
||||
) -> str:
|
||||
"""Clone a git repository into the instance directory.
|
||||
|
||||
Args:
|
||||
remote_url: Git remote URL (SSH or HTTPS)
|
||||
ssh_key_path: Path to SSH private key for authentication (optional)
|
||||
instance_dir: Path to instance directory
|
||||
branch: Branch to clone (default: main)
|
||||
|
||||
Returns:
|
||||
Path to the cloned repository
|
||||
"""
|
||||
clone_path = Path(instance_dir) / "repo-clone"
|
||||
clone_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
env = os.environ.copy()
|
||||
if ssh_key_path:
|
||||
# Use SSH key for cloning
|
||||
env["GIT_SSH_COMMAND"] = f"ssh -i {ssh_key_path} -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null"
|
||||
|
||||
cmd = [
|
||||
"git",
|
||||
"clone",
|
||||
"--branch", branch,
|
||||
"--single-branch",
|
||||
remote_url,
|
||||
str(clone_path),
|
||||
]
|
||||
|
||||
logger.debug("Cloning repository %s (branch: %s) into %s", remote_url, branch, clone_path)
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=env,
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.error("Git clone failed: %s", result.stderr)
|
||||
raise RuntimeError(f"Failed to clone repository: {result.stderr}")
|
||||
|
||||
logger.debug("Successfully cloned repository into %s", clone_path)
|
||||
return str(clone_path)
|
||||
|
||||
|
||||
def check_dirty_state(clone_path: str) -> tuple[bool, list[str]]:
|
||||
"""Check for uncommitted changes in a cloned repository.
|
||||
|
||||
Args:
|
||||
clone_path: Path to the cloned repository
|
||||
|
||||
Returns:
|
||||
Tuple of (is_dirty, list_of_changed_files)
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["git", "-C", clone_path, "status", "--short"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.warning("Failed to check git status: %s", result.stderr)
|
||||
return False, []
|
||||
|
||||
changed_files = [line.strip() for line in result.stdout.split("\n") if line.strip()]
|
||||
is_dirty = len(changed_files) > 0
|
||||
|
||||
return is_dirty, changed_files
|
||||
|
||||
|
||||
def remove_clone_directory(instance_dir: str) -> None:
|
||||
"""Remove the cloned repository from the instance directory.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
"""
|
||||
clone_path = Path(instance_dir) / "repo-clone"
|
||||
if clone_path.exists():
|
||||
import shutil
|
||||
shutil.rmtree(clone_path)
|
||||
logger.debug("Removed clone directory: %s", clone_path)
|
||||
@@ -0,0 +1,484 @@
|
||||
"""Config profile resolver service.
|
||||
|
||||
Provides recursive ordered include resolution with deterministic merge rules
|
||||
and cycle protection.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ConfigProfileCycleError(Exception):
|
||||
"""Raised when a cycle is detected in profile includes."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ConfigProfileNotFoundError(Exception):
|
||||
"""Raised when a referenced profile is not found."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolvedMount:
|
||||
"""A resolved mount with merged files and final mode."""
|
||||
|
||||
target: str
|
||||
mode: str
|
||||
files: dict[str, str] = field(default_factory=dict)
|
||||
overridden_files: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolvedProfile:
|
||||
"""The fully resolved output of a config profile."""
|
||||
|
||||
profile_id: uuid.UUID
|
||||
profile_name: str
|
||||
env_vars: dict[str, str] = field(default_factory=dict)
|
||||
runtime_hints: dict[str, Any] = field(default_factory=dict)
|
||||
mounts: dict[str, ResolvedMount] = field(default_factory=dict)
|
||||
git_mounts: list[dict[str, Any]] = field(default_factory=list)
|
||||
files: dict[str, str] = field(default_factory=dict)
|
||||
env_overrides: dict[str, str] = field(default_factory=dict)
|
||||
hint_overrides: dict[str, str] = field(default_factory=dict)
|
||||
file_overrides: dict[str, str] = field(default_factory=dict)
|
||||
mount_overrides: dict[str, str] = field(default_factory=dict)
|
||||
included_profiles: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
|
||||
def _detect_cycle(profile_id: uuid.UUID, visited: set[uuid.UUID], path: list[uuid.UUID]) -> bool:
|
||||
"""Detect if adding profile_id to path would create a cycle.
|
||||
|
||||
Args:
|
||||
profile_id: The profile ID to check.
|
||||
visited: Set of already-visited profile IDs in current resolution.
|
||||
path: Current resolution path for error reporting.
|
||||
|
||||
Returns:
|
||||
True if a cycle would be created.
|
||||
"""
|
||||
if profile_id in visited:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _merge_env_vars(
|
||||
base: dict[str, str],
|
||||
overlay: dict[str, str],
|
||||
overrides: dict[str, str],
|
||||
source_name: str,
|
||||
) -> dict[str, str]:
|
||||
"""Merge env vars, tracking overrides.
|
||||
|
||||
Later values replace earlier values.
|
||||
"""
|
||||
result = dict(base)
|
||||
for key, value in overlay.items():
|
||||
if key in result and result[key] != value:
|
||||
overrides[key] = source_name
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
def _merge_runtime_hints(
|
||||
base: dict[str, Any],
|
||||
overlay: dict[str, Any],
|
||||
overrides: dict[str, str],
|
||||
source_name: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Merge runtime hints, tracking overrides.
|
||||
|
||||
Later values replace earlier values.
|
||||
"""
|
||||
result = dict(base)
|
||||
for key, value in overlay.items():
|
||||
if key in result and result[key] != value:
|
||||
overrides[key] = source_name
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
def _merge_files(
|
||||
base: dict[str, str],
|
||||
overlay: dict[str, str],
|
||||
overrides: dict[str, str],
|
||||
source_name: str,
|
||||
) -> dict[str, str]:
|
||||
"""Merge file maps, tracking overrides.
|
||||
|
||||
Later relative file paths win.
|
||||
"""
|
||||
result = dict(base)
|
||||
for path, content in overlay.items():
|
||||
if path in result and result[path] != content:
|
||||
overrides[path] = source_name
|
||||
result[path] = content
|
||||
return result
|
||||
|
||||
|
||||
def _merge_mounts(
|
||||
base: dict[str, ResolvedMount],
|
||||
overlay: list[dict[str, Any]],
|
||||
overrides: dict[str, str],
|
||||
source_name: str,
|
||||
) -> dict[str, ResolvedMount]:
|
||||
"""Merge mounts, tracking overrides.
|
||||
|
||||
Mounts with the same target path have their file maps merged and later
|
||||
relative file paths win. Mode conflicts: later layer wins.
|
||||
"""
|
||||
result = dict(base)
|
||||
for mount_data in overlay:
|
||||
target = mount_data["target"]
|
||||
mode = mount_data.get("mode", "rw")
|
||||
files = mount_data.get("files", {})
|
||||
|
||||
if target in result:
|
||||
existing = result[target]
|
||||
merged_files = dict(existing.files)
|
||||
file_overrides = dict(existing.overridden_files)
|
||||
for rel_path, content in files.items():
|
||||
if rel_path in merged_files and merged_files[rel_path] != content:
|
||||
file_overrides[rel_path] = source_name
|
||||
merged_files[rel_path] = content
|
||||
if existing.mode != mode:
|
||||
overrides[target] = source_name
|
||||
result[target] = ResolvedMount(
|
||||
target=target,
|
||||
mode=mode,
|
||||
files=merged_files,
|
||||
overridden_files=file_overrides,
|
||||
)
|
||||
else:
|
||||
result[target] = ResolvedMount(
|
||||
target=target,
|
||||
mode=mode,
|
||||
files=dict(files),
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _merge_git_mounts(
|
||||
base: list[dict[str, Any]],
|
||||
overlay: list[dict[str, Any]],
|
||||
source_name: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Merge git mounts from included profiles.
|
||||
|
||||
Later mounts override earlier ones with the same remote_url + target_path combo.
|
||||
"""
|
||||
result = list(base)
|
||||
# Build lookup by (remote_url, target_path)
|
||||
seen = {(m["remote_url"], m["target_path"]): i for i, m in enumerate(result)}
|
||||
for mount in overlay:
|
||||
key = (mount["remote_url"], mount["target_path"])
|
||||
if key in seen:
|
||||
result[seen[key]] = dict(mount)
|
||||
else:
|
||||
seen[key] = len(result)
|
||||
result.append(dict(mount))
|
||||
return result
|
||||
|
||||
|
||||
async def _resolve_profile_recursive(
|
||||
session: AsyncSession,
|
||||
profile_id: uuid.UUID,
|
||||
visited: set[uuid.UUID],
|
||||
path: list[uuid.UUID],
|
||||
) -> ResolvedProfile:
|
||||
"""Recursively resolve a profile and its includes.
|
||||
|
||||
Args:
|
||||
session: Database session.
|
||||
profile_id: Profile ID to resolve.
|
||||
visited: Set of already-visited profile IDs in current resolution chain.
|
||||
path: Current resolution path for error reporting.
|
||||
|
||||
Returns:
|
||||
ResolvedProfile with all includes merged.
|
||||
|
||||
Raises:
|
||||
ConfigProfileCycleError: If a cycle is detected.
|
||||
ConfigProfileNotFoundError: If the profile is not found.
|
||||
"""
|
||||
if _detect_cycle(profile_id, visited, path):
|
||||
cycle_path = " -> ".join(str(p) for p in path + [profile_id])
|
||||
raise ConfigProfileCycleError(f"Cycle detected in profile includes: {cycle_path}")
|
||||
|
||||
profile = await session.get(ConfigProfile, profile_id)
|
||||
if profile is None:
|
||||
raise ConfigProfileNotFoundError(f"Config profile not found: {profile_id}")
|
||||
|
||||
new_visited = visited | {profile_id}
|
||||
new_path = path + [profile_id]
|
||||
|
||||
result = ResolvedProfile(
|
||||
profile_id=profile.id,
|
||||
profile_name=profile.name,
|
||||
)
|
||||
|
||||
# Resolve includes in order
|
||||
include_query = (
|
||||
select(ConfigProfileInclude)
|
||||
.where(ConfigProfileInclude.profile_id == profile_id)
|
||||
.order_by(ConfigProfileInclude.order_index)
|
||||
)
|
||||
include_result = await session.execute(include_query)
|
||||
includes = include_result.scalars().all()
|
||||
|
||||
for include in includes:
|
||||
included = await _resolve_profile_recursive(
|
||||
session, include.included_profile_id, new_visited, new_path
|
||||
)
|
||||
result.included_profiles.append({
|
||||
"id": str(included.profile_id),
|
||||
"name": included.profile_name,
|
||||
})
|
||||
|
||||
result.env_vars = _merge_env_vars(
|
||||
result.env_vars, included.env_vars, result.env_overrides, included.profile_name
|
||||
)
|
||||
result.runtime_hints = _merge_runtime_hints(
|
||||
result.runtime_hints,
|
||||
included.runtime_hints,
|
||||
result.hint_overrides,
|
||||
included.profile_name,
|
||||
)
|
||||
result.files = _merge_files(
|
||||
result.files, included.files, result.file_overrides, included.profile_name
|
||||
)
|
||||
result.mounts = _merge_mounts(
|
||||
result.mounts,
|
||||
[
|
||||
{"target": m.target, "mode": m.mode, "files": m.files}
|
||||
for m in included.mounts.values()
|
||||
],
|
||||
result.mount_overrides,
|
||||
included.profile_name,
|
||||
)
|
||||
result.git_mounts = _merge_git_mounts(
|
||||
result.git_mounts, included.git_mounts, included.profile_name
|
||||
)
|
||||
|
||||
# Apply the profile's own settings (selected profile overrides includes)
|
||||
result.env_vars = _merge_env_vars(
|
||||
result.env_vars,
|
||||
profile.env_vars or {},
|
||||
result.env_overrides,
|
||||
profile.name,
|
||||
)
|
||||
result.runtime_hints = _merge_runtime_hints(
|
||||
result.runtime_hints,
|
||||
profile.runtime_hints or {},
|
||||
result.hint_overrides,
|
||||
profile.name,
|
||||
)
|
||||
result.files = _merge_files(
|
||||
result.files,
|
||||
profile.files or {},
|
||||
result.file_overrides,
|
||||
profile.name,
|
||||
)
|
||||
result.mounts = _merge_mounts(
|
||||
result.mounts,
|
||||
profile.mounts or [],
|
||||
result.mount_overrides,
|
||||
profile.name,
|
||||
)
|
||||
result.git_mounts = _merge_git_mounts(
|
||||
result.git_mounts,
|
||||
profile.git_mounts or [],
|
||||
profile.name,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def resolve_profile(
|
||||
session: AsyncSession,
|
||||
profile_id: uuid.UUID,
|
||||
) -> ResolvedProfile:
|
||||
"""Resolve a config profile with all includes.
|
||||
|
||||
Args:
|
||||
session: Database session.
|
||||
profile_id: Profile ID to resolve.
|
||||
|
||||
Returns:
|
||||
ResolvedProfile with merged env vars, runtime hints, mounts, and files.
|
||||
|
||||
Raises:
|
||||
ConfigProfileCycleError: If a cycle is detected in includes.
|
||||
ConfigProfileNotFoundError: If the profile is not found.
|
||||
"""
|
||||
return await _resolve_profile_recursive(session, profile_id, set(), [])
|
||||
|
||||
|
||||
async def check_include_cycle(
|
||||
session: AsyncSession,
|
||||
profile_id: uuid.UUID,
|
||||
new_include_id: uuid.UUID | None = None,
|
||||
) -> list[uuid.UUID] | None:
|
||||
"""Check if adding an include would create a cycle.
|
||||
|
||||
Used at save time to validate include relationships before persisting.
|
||||
|
||||
Args:
|
||||
session: Database session.
|
||||
profile_id: The profile that would receive the new include.
|
||||
new_include_id: Optional new profile to include. If None, checks existing includes.
|
||||
|
||||
Returns:
|
||||
The cycle path as a list of UUIDs if a cycle exists, otherwise None.
|
||||
"""
|
||||
|
||||
async def _check_from(
|
||||
current_id: uuid.UUID,
|
||||
target_id: uuid.UUID,
|
||||
visited: set[uuid.UUID],
|
||||
path: list[uuid.UUID],
|
||||
) -> list[uuid.UUID] | None:
|
||||
if current_id in visited:
|
||||
if current_id == target_id:
|
||||
return path + [current_id]
|
||||
return None
|
||||
if current_id == target_id and path:
|
||||
return path + [current_id]
|
||||
|
||||
new_visited = visited | {current_id}
|
||||
new_path = path + [current_id]
|
||||
|
||||
include_query = (
|
||||
select(ConfigProfileInclude)
|
||||
.where(ConfigProfileInclude.profile_id == current_id)
|
||||
.order_by(ConfigProfileInclude.order_index)
|
||||
)
|
||||
include_result = await session.execute(include_query)
|
||||
includes = include_result.scalars().all()
|
||||
|
||||
for include in includes:
|
||||
cycle = await _check_from(
|
||||
include.included_profile_id, target_id, new_visited, new_path
|
||||
)
|
||||
if cycle is not None:
|
||||
return cycle
|
||||
return None
|
||||
|
||||
# Check if new_include_id can reach profile_id (would create cycle)
|
||||
if new_include_id is not None:
|
||||
cycle = await _check_from(new_include_id, profile_id, set(), [])
|
||||
if cycle is not None:
|
||||
return cycle
|
||||
|
||||
# Also check existing includes for cycles
|
||||
cycle = await _check_from(profile_id, profile_id, set(), [])
|
||||
if cycle is not None and len(cycle) > 1:
|
||||
return cycle
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def apply_resolved_profile(
|
||||
instance_dir: str,
|
||||
resolved: ResolvedProfile,
|
||||
) -> tuple[dict[str, str], dict[str, str], list[dict], dict[str, Any]]:
|
||||
"""Apply a resolved profile to an instance directory.
|
||||
|
||||
Stages files, writes env vars, and prepares mount volumes.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to the instance directory.
|
||||
resolved: The resolved profile.
|
||||
|
||||
Returns:
|
||||
Tuple of (env_vars, files, volume_mounts, runtime_hints).
|
||||
env_vars: Merged environment variables.
|
||||
files: Relative file paths to content for the instance.
|
||||
volume_mounts: List of Docker volume mount dicts.
|
||||
runtime_hints: Extracted runtime hints.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
instance_path = Path(instance_dir)
|
||||
env_vars = dict(resolved.env_vars)
|
||||
files = dict(resolved.files)
|
||||
volume_mounts = []
|
||||
|
||||
# Write profile files to instance directory
|
||||
for file_path, content in files.items():
|
||||
full_path = instance_path / file_path
|
||||
try:
|
||||
full_path.resolve().relative_to(instance_path.resolve())
|
||||
except ValueError:
|
||||
logger.warning("Profile file path escapes instance directory: %s", file_path)
|
||||
continue
|
||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
full_path.write_text(content)
|
||||
|
||||
# Stage mount files and prepare volume mounts
|
||||
for mount in resolved.mounts.values():
|
||||
mount_dir = instance_path / "mounts" / mount.target.lstrip("/").replace("/", "_")
|
||||
mount_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for file_path, content in mount.files.items():
|
||||
full_path = mount_dir / file_path
|
||||
try:
|
||||
full_path.resolve().relative_to(mount_dir.resolve())
|
||||
except ValueError:
|
||||
logger.warning("Mount file path escapes mount directory: %s", file_path)
|
||||
continue
|
||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
full_path.write_text(content)
|
||||
|
||||
volume_mounts.append({
|
||||
"source": str(mount_dir),
|
||||
"target": mount.target,
|
||||
"type": "bind",
|
||||
})
|
||||
|
||||
return env_vars, files, volume_mounts, resolved.runtime_hints
|
||||
|
||||
|
||||
def resolved_profile_to_dict(resolved: ResolvedProfile) -> dict[str, Any]:
|
||||
"""Convert a ResolvedProfile to a plain dict for serialization.
|
||||
|
||||
Args:
|
||||
resolved: The resolved profile.
|
||||
|
||||
Returns:
|
||||
Dict with env_vars, runtime_hints, mounts, files, and metadata.
|
||||
"""
|
||||
return {
|
||||
"profile_id": str(resolved.profile_id),
|
||||
"profile_name": resolved.profile_name,
|
||||
"env_vars": resolved.env_vars,
|
||||
"runtime_hints": resolved.runtime_hints,
|
||||
"mounts": [
|
||||
{
|
||||
"target": m.target,
|
||||
"mode": m.mode,
|
||||
"files": m.files,
|
||||
"overridden_files": m.overridden_files,
|
||||
}
|
||||
for m in resolved.mounts.values()
|
||||
],
|
||||
"files": resolved.files,
|
||||
"overrides": {
|
||||
"env_vars": resolved.env_overrides,
|
||||
"runtime_hints": resolved.hint_overrides,
|
||||
"files": resolved.file_overrides,
|
||||
"mounts": resolved.mount_overrides,
|
||||
},
|
||||
"git_mounts": resolved.git_mounts,
|
||||
"included_profiles": resolved.included_profiles,
|
||||
}
|
||||
@@ -0,0 +1,549 @@
|
||||
"""Docker service for managing tool instances."""
|
||||
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def render_compose_template(template: str, variables: dict[str, Any]) -> str:
|
||||
"""Render a Docker Compose template with variable substitution.
|
||||
|
||||
Args:
|
||||
template: The compose template string
|
||||
variables: Dictionary of variable names to values
|
||||
|
||||
Returns:
|
||||
Rendered compose file content
|
||||
"""
|
||||
result = template
|
||||
for key, value in variables.items():
|
||||
placeholder = f"{{{{{key}}}}}"
|
||||
result = result.replace(placeholder, str(value))
|
||||
return result
|
||||
|
||||
|
||||
def ensure_instance_directory(instance_id: str, base_path: str | None = None) -> str:
|
||||
"""Create and return the instance directory path.
|
||||
|
||||
Args:
|
||||
instance_id: Unique instance identifier
|
||||
base_path: Base directory for all instances (defaults to Settings.instance_base_path)
|
||||
|
||||
Returns:
|
||||
Absolute path to instance directory
|
||||
"""
|
||||
if base_path is None:
|
||||
from src.config import Settings
|
||||
|
||||
base_path = Settings().instance_base_path
|
||||
instance_dir = Path(base_path) / instance_id
|
||||
instance_dir.mkdir(parents=True, exist_ok=True)
|
||||
return str(instance_dir.absolute())
|
||||
|
||||
|
||||
def write_compose_file(instance_dir: str, content: str) -> str:
|
||||
"""Write the rendered compose file to the instance directory.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
content: Rendered compose content
|
||||
|
||||
Returns:
|
||||
Path to the compose file
|
||||
"""
|
||||
compose_path = Path(instance_dir) / "docker-compose.yml"
|
||||
compose_path.write_text(content)
|
||||
return str(compose_path)
|
||||
|
||||
|
||||
def write_env_file(instance_dir: str, env_vars: dict[str, str]) -> str:
|
||||
"""Write environment variables to a .env file.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
env_vars: Dictionary of env var names to values
|
||||
|
||||
Returns:
|
||||
Path to the env file
|
||||
"""
|
||||
env_path = Path(instance_dir) / ".env"
|
||||
lines = [f'{key}="{value}"' for key, value in env_vars.items()]
|
||||
env_path.write_text("\n".join(lines) + "\n")
|
||||
return str(env_path)
|
||||
|
||||
|
||||
def write_config_files(instance_dir: str, files: dict[str, str]) -> None:
|
||||
"""Write config files to the instance directory.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
files: Dictionary of file paths (relative to instance dir) to content
|
||||
"""
|
||||
instance_path = Path(instance_dir)
|
||||
for file_path, content in files.items():
|
||||
# Ensure the path is within the instance directory (security)
|
||||
full_path = instance_path / file_path
|
||||
try:
|
||||
full_path.resolve().relative_to(instance_path.resolve())
|
||||
except ValueError:
|
||||
raise ValueError(f"File path '{file_path}' escapes instance directory")
|
||||
|
||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
full_path.write_text(content)
|
||||
|
||||
|
||||
def execute_compose_command(
|
||||
compose_path: str, action: str, timeout: int = 60, env_file: str | None = None
|
||||
) -> tuple[int, str, str]:
|
||||
"""Execute a docker compose command.
|
||||
|
||||
Args:
|
||||
compose_path: Path to docker-compose.yml
|
||||
action: The compose action (up, down, start, stop, restart)
|
||||
timeout: Command timeout in seconds
|
||||
env_file: Optional path to .env file for environment variables
|
||||
|
||||
Returns:
|
||||
Tuple of (returncode, stdout, stderr)
|
||||
"""
|
||||
instance_dir = Path(compose_path).parent
|
||||
|
||||
cmd = ["docker", "compose", "-f", compose_path]
|
||||
|
||||
if env_file:
|
||||
cmd.extend(["--env-file", env_file])
|
||||
|
||||
if action == "up":
|
||||
cmd.extend(["up", "-d"])
|
||||
elif action == "down":
|
||||
cmd.extend(["down", "-v"])
|
||||
elif action in ("start", "stop", "restart"):
|
||||
cmd.append(action)
|
||||
else:
|
||||
raise ValueError(f"Unknown compose action: {action}")
|
||||
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
cwd=str(instance_dir),
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return result.returncode, result.stdout, result.stderr
|
||||
|
||||
|
||||
def get_container_id(instance_name: str) -> str | None:
|
||||
"""Get the container ID for a compose service.
|
||||
|
||||
Searches all containers including stopped/exited ones.
|
||||
|
||||
Args:
|
||||
instance_name: The service name in compose
|
||||
|
||||
Returns:
|
||||
Container ID or None if not found
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["docker", "ps", "-a", "-q", "--filter", f"name={instance_name}"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return result.stdout.strip().split("\n")[0]
|
||||
return None
|
||||
|
||||
|
||||
def get_container_name(instance_name: str) -> str | None:
|
||||
"""Get the full container name for a compose service.
|
||||
|
||||
Searches all containers including stopped/exited ones.
|
||||
|
||||
Args:
|
||||
instance_name: The service name in compose
|
||||
|
||||
Returns:
|
||||
Container name or None if not found
|
||||
"""
|
||||
result = subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"ps",
|
||||
"-a",
|
||||
"--format",
|
||||
"{{.Names}}",
|
||||
"--filter",
|
||||
f"name={instance_name}",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return result.stdout.strip().split("\n")[0]
|
||||
return None
|
||||
|
||||
|
||||
def connect_container_to_network(
|
||||
container_name: str, network_name: str = "backend"
|
||||
) -> bool:
|
||||
"""Connect a Docker container to an existing network.
|
||||
|
||||
Args:
|
||||
container_name: Name or ID of the container
|
||||
network_name: Name of the Docker network (default: backend)
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["docker", "network", "connect", network_name, container_name],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
return result.returncode == 0
|
||||
|
||||
|
||||
def get_container_status(container_id: str) -> dict[str, Any]:
|
||||
"""Get the status of a Docker container.
|
||||
|
||||
Args:
|
||||
container_id: Docker container ID
|
||||
|
||||
Returns:
|
||||
Dict with 'status' (running, exited, restarting, not_found),
|
||||
'exit_code' (int or None), and 'health' (health status or None)
|
||||
"""
|
||||
result = subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"inspect",
|
||||
"-f",
|
||||
"{{.State.Status}}|{{.State.ExitCode}}|{{if .State.Health}}{{.State.Health.Status}}{{else}}none{{end}}",
|
||||
container_id,
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
return {"status": "not_found", "exit_code": None, "health": None}
|
||||
|
||||
parts = result.stdout.strip().split("|")
|
||||
status = parts[0] if parts else "unknown"
|
||||
exit_code = int(parts[1]) if len(parts) > 1 and parts[1].isdigit() else None
|
||||
health = parts[2] if len(parts) > 2 and parts[2] != "none" else None
|
||||
|
||||
return {"status": status, "exit_code": exit_code, "health": health}
|
||||
|
||||
|
||||
def wait_for_container_running(
|
||||
container_id: str, timeout: int = 30, interval: float = 2.0
|
||||
) -> dict[str, Any]:
|
||||
"""Wait for a container to reach the running state.
|
||||
|
||||
Polls docker inspect until the container status is "running" or timeout.
|
||||
|
||||
Args:
|
||||
container_id: Docker container ID
|
||||
timeout: Maximum seconds to wait
|
||||
interval: Seconds between polls
|
||||
|
||||
Returns:
|
||||
Dict with 'success' (bool), 'status' (str), 'exit_code' (int or None),
|
||||
and 'waited_seconds' (float)
|
||||
"""
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
info = get_container_status(container_id)
|
||||
|
||||
if info["status"] == "running":
|
||||
return {
|
||||
"success": True,
|
||||
"status": "running",
|
||||
"exit_code": None,
|
||||
"waited_seconds": time.time() - start_time,
|
||||
}
|
||||
|
||||
if info["status"] == "exited":
|
||||
return {
|
||||
"success": False,
|
||||
"status": "exited",
|
||||
"exit_code": info["exit_code"],
|
||||
"waited_seconds": time.time() - start_time,
|
||||
}
|
||||
|
||||
if info["status"] == "not_found":
|
||||
return {
|
||||
"success": False,
|
||||
"status": "not_found",
|
||||
"exit_code": None,
|
||||
"waited_seconds": time.time() - start_time,
|
||||
}
|
||||
|
||||
time.sleep(interval)
|
||||
|
||||
# Timeout reached
|
||||
info = get_container_status(container_id)
|
||||
return {
|
||||
"success": False,
|
||||
"status": info["status"],
|
||||
"exit_code": info["exit_code"],
|
||||
"waited_seconds": time.time() - start_time,
|
||||
}
|
||||
|
||||
|
||||
def get_container_logs(container_id: str, tail: int = 100) -> str:
|
||||
"""Get the logs of a Docker container.
|
||||
|
||||
Args:
|
||||
container_id: Docker container ID
|
||||
tail: Number of lines to return
|
||||
|
||||
Returns:
|
||||
Container logs
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["docker", "logs", "--tail", str(tail), container_id],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
return result.stdout
|
||||
return f"Failed to get logs: {result.stderr}"
|
||||
|
||||
|
||||
def find_free_port(start: int = 10000, end: int = 20000) -> int:
|
||||
"""Find a free TCP port in the given range.
|
||||
|
||||
Args:
|
||||
start: Start of port range
|
||||
end: End of port range
|
||||
|
||||
Returns:
|
||||
Free port number
|
||||
"""
|
||||
import socket
|
||||
|
||||
for port in range(start, end):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
if s.connect_ex(("localhost", port)) != 0:
|
||||
return port
|
||||
|
||||
raise RuntimeError(f"No free port found in range {start}-{end}")
|
||||
|
||||
|
||||
def start_cloudflared_tunnel(
|
||||
container_name: str, port: int, timeout: int = 30
|
||||
) -> dict[str, str]:
|
||||
"""Start a temporary Cloudflare tunnel for a container.
|
||||
|
||||
Uses 'cloudflared tunnel --url' to create a temporary tunnel
|
||||
with a random trycloudflare.com URL.
|
||||
|
||||
Args:
|
||||
container_name: Name of the Docker container to tunnel to
|
||||
port: Port number the container listens on
|
||||
timeout: Maximum seconds to wait for tunnel URL
|
||||
|
||||
Returns:
|
||||
Dict with 'url' (the public tunnel URL) and 'pid' (process ID)
|
||||
"""
|
||||
import subprocess
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# First verify the container is accessible
|
||||
logger.info("Checking connectivity to %s:%d...", container_name, port)
|
||||
for attempt in range(10):
|
||||
check = subprocess.run(
|
||||
[
|
||||
"curl",
|
||||
"-s",
|
||||
"-o",
|
||||
"/dev/null",
|
||||
"-w",
|
||||
"%{http_code}",
|
||||
f"http://{container_name}:{port}",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
)
|
||||
logger.info(
|
||||
"Connectivity check %d: http_code=%s", attempt + 1, check.stdout.strip()
|
||||
)
|
||||
if check.returncode == 0:
|
||||
break
|
||||
time.sleep(1)
|
||||
else:
|
||||
logger.warning(
|
||||
"Container %s:%d not responding to curl checks", container_name, port
|
||||
)
|
||||
|
||||
# Run cloudflared in background, capture output
|
||||
logger.info("Starting cloudflared tunnel to http://%s:%d", container_name, port)
|
||||
proc = subprocess.Popen(
|
||||
["cloudflared", "tunnel", "--url", f"http://{container_name}:{port}"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
)
|
||||
|
||||
# Wait for the URL to appear in output
|
||||
url_pattern = re.compile(r"https://[a-z0-9-]+\.trycloudflare\.com")
|
||||
start_time = time.time()
|
||||
url = None
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
# Read available output
|
||||
import select
|
||||
|
||||
readable, _, _ = select.select([proc.stdout], [], [], 1.0)
|
||||
if readable:
|
||||
line = proc.stdout.readline()
|
||||
if line:
|
||||
match = url_pattern.search(line)
|
||||
if match:
|
||||
url = match.group(0)
|
||||
break
|
||||
|
||||
if not url:
|
||||
proc.terminate()
|
||||
proc.wait(timeout=5)
|
||||
raise RuntimeError(
|
||||
f"Failed to get tunnel URL within {timeout}s. "
|
||||
f"cloudflared output may contain errors."
|
||||
)
|
||||
|
||||
return {"url": url, "pid": str(proc.pid)}
|
||||
|
||||
|
||||
def stop_cloudflared_tunnel(pid: str) -> None:
|
||||
"""Stop a cloudflared tunnel process.
|
||||
|
||||
Args:
|
||||
pid: Process ID of the cloudflared tunnel
|
||||
"""
|
||||
import signal
|
||||
|
||||
try:
|
||||
os.kill(int(pid), signal.SIGTERM)
|
||||
except ProcessLookupError:
|
||||
pass # Already stopped
|
||||
|
||||
|
||||
def recreate_tunnel(
|
||||
container_name: str, port: int, old_pid: str | None = None
|
||||
) -> dict[str, str]:
|
||||
"""Recreate a temporary Cloudflare tunnel.
|
||||
|
||||
Stops the old tunnel (if pid provided) and starts a new one.
|
||||
|
||||
Args:
|
||||
container_name: Name of the Docker container to tunnel to
|
||||
port: Port number the container listens on
|
||||
old_pid: Optional PID of the old tunnel process to stop
|
||||
|
||||
Returns:
|
||||
Dict with 'url' and 'pid' for the new tunnel
|
||||
"""
|
||||
if old_pid:
|
||||
stop_cloudflared_tunnel(old_pid)
|
||||
|
||||
return start_cloudflared_tunnel(container_name, port)
|
||||
|
||||
|
||||
def check_tunnel_health(url: str, timeout: int = 10) -> dict[str, Any]:
|
||||
"""Check if a tunnel URL is healthy with smart error classification.
|
||||
|
||||
Args:
|
||||
url: The tunnel URL to check
|
||||
timeout: Request timeout in seconds
|
||||
|
||||
Returns:
|
||||
Dict with 'tunnel_status' (healthy, unreachable, error_response, not_applicable),
|
||||
'status_code' (int or None), 'healthy' (bool), and 'error' (str or None)
|
||||
"""
|
||||
import subprocess
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
"curl",
|
||||
"-s",
|
||||
"-o",
|
||||
"/dev/null",
|
||||
"-w",
|
||||
"%{http_code}",
|
||||
"--max-time",
|
||||
str(timeout),
|
||||
url,
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout + 5,
|
||||
)
|
||||
status_code = int(result.stdout.strip())
|
||||
|
||||
if 200 <= status_code < 400:
|
||||
return {
|
||||
"tunnel_status": "healthy",
|
||||
"status_code": status_code,
|
||||
"healthy": True,
|
||||
"error": None,
|
||||
}
|
||||
elif status_code in (502, 503, 504):
|
||||
# Application error, not tunnel error
|
||||
return {
|
||||
"tunnel_status": "error_response",
|
||||
"status_code": status_code,
|
||||
"healthy": False,
|
||||
"error": f"Application returned HTTP {status_code}",
|
||||
}
|
||||
else:
|
||||
return {
|
||||
"tunnel_status": "error_response",
|
||||
"status_code": status_code,
|
||||
"healthy": False,
|
||||
"error": f"HTTP {status_code}",
|
||||
}
|
||||
except subprocess.TimeoutExpired:
|
||||
return {
|
||||
"tunnel_status": "unreachable",
|
||||
"status_code": None,
|
||||
"healthy": False,
|
||||
"error": "Tunnel request timed out",
|
||||
}
|
||||
except (ValueError, Exception) as e:
|
||||
error_str = str(e).lower()
|
||||
# Classify connection errors
|
||||
if any(
|
||||
err in error_str
|
||||
for err in [
|
||||
"connection refused",
|
||||
"econnrefused",
|
||||
"could not resolve",
|
||||
"nodename",
|
||||
]
|
||||
):
|
||||
return {
|
||||
"tunnel_status": "unreachable",
|
||||
"status_code": None,
|
||||
"healthy": False,
|
||||
"error": f"Tunnel unreachable: {e}",
|
||||
}
|
||||
return {
|
||||
"tunnel_status": "unreachable",
|
||||
"status_code": None,
|
||||
"healthy": False,
|
||||
"error": str(e),
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Docker build service for building images from Dockerfiles."""
|
||||
|
||||
import logging
|
||||
import subprocess
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def build_image(instance_dir: str, dockerfile: str, tag: str, build_context: dict | None = None) -> tuple[int, str, str]:
|
||||
"""Build a Docker image from a Dockerfile.
|
||||
|
||||
Args:
|
||||
instance_dir: Directory containing the Dockerfile
|
||||
dockerfile: Dockerfile content
|
||||
tag: Image tag to apply
|
||||
build_context: Optional build context files {path: content}
|
||||
|
||||
Returns:
|
||||
Tuple of (returncode, stdout, stderr)
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
# Write Dockerfile
|
||||
dockerfile_path = Path(instance_dir) / "Dockerfile"
|
||||
dockerfile_path.write_text(dockerfile)
|
||||
logger.debug("Wrote Dockerfile to %s", dockerfile_path)
|
||||
|
||||
# Write build context files
|
||||
if build_context:
|
||||
for file_path, content in build_context.items():
|
||||
full_path = Path(instance_dir) / file_path
|
||||
# Security: ensure path doesn't escape instance_dir
|
||||
try:
|
||||
full_path.resolve().relative_to(Path(instance_dir).resolve())
|
||||
except ValueError:
|
||||
logger.error("Build context file path escapes instance directory: %s", file_path)
|
||||
raise ValueError(f"Build context file path '{file_path}' escapes instance directory")
|
||||
|
||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
full_path.write_text(content)
|
||||
logger.debug("Wrote build context file: %s", full_path)
|
||||
|
||||
# Build image
|
||||
logger.debug("Building Docker image with tag: %s", tag)
|
||||
cmd = [
|
||||
"docker", "build",
|
||||
"-t", tag,
|
||||
"-f", str(dockerfile_path),
|
||||
instance_dir,
|
||||
]
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300, # 5 minute timeout for builds
|
||||
)
|
||||
logger.debug("Docker build completed: returncode=%d", result.returncode)
|
||||
if result.returncode != 0:
|
||||
logger.error("Docker build failed: %s", result.stderr[:1000])
|
||||
return result.returncode, result.stdout, result.stderr
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.error("Docker build timed out after 300 seconds")
|
||||
return 1, "", "Build timed out after 300 seconds"
|
||||
except Exception as exc:
|
||||
logger.exception("Docker build failed: %s", exc)
|
||||
return 1, "", str(exc)
|
||||
@@ -0,0 +1,66 @@
|
||||
"""Readiness probe service for checking if containers are ready."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import subprocess
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def execute_probe(
|
||||
container_id: str,
|
||||
command: str,
|
||||
timeout: int = 30,
|
||||
interval: int = 2,
|
||||
) -> tuple[bool, list[str]]:
|
||||
"""Execute a readiness probe command inside a container.
|
||||
|
||||
Args:
|
||||
container_id: Docker container ID or name
|
||||
command: Command to execute inside the container
|
||||
timeout: Maximum total time to wait (seconds)
|
||||
interval: Time between retries (seconds)
|
||||
|
||||
Returns:
|
||||
Tuple of (success, logs)
|
||||
"""
|
||||
logs = []
|
||||
start_time = asyncio.get_event_loop().time()
|
||||
attempt = 0
|
||||
|
||||
while True:
|
||||
attempt += 1
|
||||
elapsed = asyncio.get_event_loop().time() - start_time
|
||||
|
||||
if elapsed >= timeout:
|
||||
logs.append(f"Probe timed out after {timeout}s ({attempt} attempts)")
|
||||
return False, logs
|
||||
|
||||
try:
|
||||
logger.debug("Probe attempt %d: %s", attempt, command)
|
||||
|
||||
# Execute command inside container
|
||||
result = subprocess.run(
|
||||
["docker", "exec", container_id, "sh", "-c", command],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=interval, # Each attempt has its own timeout
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
logs.append(f"Attempt {attempt}: Success")
|
||||
if result.stdout:
|
||||
logs.append(f"Output: {result.stdout.strip()}")
|
||||
return True, logs
|
||||
else:
|
||||
logs.append(f"Attempt {attempt}: Failed (exit code {result.returncode})")
|
||||
if result.stderr:
|
||||
logs.append(f"Stderr: {result.stderr.strip()[:200]}")
|
||||
|
||||
except subprocess.TimeoutExpired:
|
||||
logs.append(f"Attempt {attempt}: Command timed out")
|
||||
except Exception as exc:
|
||||
logs.append(f"Attempt {attempt}: Error - {exc}")
|
||||
|
||||
# Wait before next attempt
|
||||
await asyncio.sleep(interval)
|
||||
@@ -0,0 +1,73 @@
|
||||
"""SSH key service utilities for preparing keys for container use."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def _get_fernet() -> Fernet:
|
||||
"""Generate a valid Fernet key from the session secret."""
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
settings = Settings()
|
||||
key_bytes = hashlib.sha256(settings.session_secret.encode()).digest()
|
||||
key = base64.urlsafe_b64encode(key_bytes)
|
||||
return Fernet(key)
|
||||
|
||||
|
||||
def prepare_ssh_key_files(instance_dir: str, ssh_key) -> str:
|
||||
"""Decrypt and write SSH key files to instance directory for container mounting.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
ssh_key: SSHKey model instance with encrypted private key
|
||||
|
||||
Returns:
|
||||
Path to the .ssh directory
|
||||
"""
|
||||
ssh_dir = Path(instance_dir) / ".ssh"
|
||||
ssh_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Decrypt private key
|
||||
fernet = _get_fernet()
|
||||
private_key = fernet.decrypt(ssh_key.private_key_encrypted.encode()).decode()
|
||||
|
||||
# Write private key with restricted permissions
|
||||
private_key_path = ssh_dir / "id_ed25519"
|
||||
private_key_path.write_text(private_key)
|
||||
os.chmod(private_key_path, 0o600)
|
||||
|
||||
# Write public key
|
||||
public_key_path = ssh_dir / "id_ed25519.pub"
|
||||
public_key_path.write_text(ssh_key.public_key)
|
||||
os.chmod(public_key_path, 0o644)
|
||||
|
||||
# Write SSH config
|
||||
config_path = ssh_dir / "config"
|
||||
config_content = """Host *
|
||||
StrictHostKeyChecking no
|
||||
UserKnownHostsFile /dev/null
|
||||
IdentityFile ~/.ssh/id_ed25519
|
||||
IdentitiesOnly yes
|
||||
"""
|
||||
config_path.write_text(config_content)
|
||||
os.chmod(config_path, 0o644)
|
||||
|
||||
return str(ssh_dir)
|
||||
|
||||
|
||||
def cleanup_ssh_key_files(instance_dir: str) -> None:
|
||||
"""Remove temporary SSH key files from instance directory.
|
||||
|
||||
Args:
|
||||
instance_dir: Path to instance directory
|
||||
"""
|
||||
ssh_dir = Path(instance_dir) / ".ssh"
|
||||
if ssh_dir.exists():
|
||||
for file_path in ssh_dir.iterdir():
|
||||
file_path.unlink()
|
||||
ssh_dir.rmdir()
|
||||
@@ -0,0 +1,161 @@
|
||||
"""Terminal session manager for WebSocket connections."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
from fastapi import WebSocket
|
||||
|
||||
from src.services.terminal_session import TerminalSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TerminalManager:
|
||||
"""Manages active terminal sessions with persistence support."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
# Track sessions by instance_id for persistence
|
||||
self._sessions: dict[str, TerminalSession] = {}
|
||||
self._idle_check_task: asyncio.Task | None = None
|
||||
self._start_idle_check()
|
||||
|
||||
def _start_idle_check(self) -> None:
|
||||
"""Start the idle timeout background task."""
|
||||
if self._idle_check_task is not None and not self._idle_check_task.done():
|
||||
return
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
self._idle_check_task = loop.create_task(self._idle_check_loop())
|
||||
except RuntimeError:
|
||||
# No event loop running yet, will be started lazily
|
||||
pass
|
||||
|
||||
async def _idle_check_loop(self) -> None:
|
||||
"""Periodically check for idle sessions and clean them up."""
|
||||
while True:
|
||||
try:
|
||||
await asyncio.sleep(60) # Check every minute
|
||||
await self._cleanup_idle_sessions()
|
||||
except Exception as exc:
|
||||
logger.error("Error in idle check loop: %s", exc)
|
||||
|
||||
async def _cleanup_idle_sessions(self) -> None:
|
||||
"""Clean up sessions that have been idle for too long."""
|
||||
idle_sessions = []
|
||||
for instance_id, session in list(self._sessions.items()):
|
||||
if session.is_idle():
|
||||
idle_sessions.append(instance_id)
|
||||
|
||||
for instance_id in idle_sessions:
|
||||
logger.info("Cleaning up idle terminal session for instance %s", instance_id)
|
||||
session = self._sessions.pop(instance_id, None)
|
||||
if session:
|
||||
await session.close()
|
||||
|
||||
async def get_or_create_session(
|
||||
self,
|
||||
instance_id: uuid.UUID,
|
||||
container_id: str,
|
||||
startup_command: str | None = None,
|
||||
) -> TerminalSession:
|
||||
"""Get existing session or create a new one."""
|
||||
# Ensure idle check is running (lazy start)
|
||||
self._start_idle_check()
|
||||
|
||||
instance_id_str = str(instance_id)
|
||||
|
||||
# Check for existing session
|
||||
if instance_id_str in self._sessions:
|
||||
session = self._sessions[instance_id_str]
|
||||
|
||||
# Check if session is still alive
|
||||
if session.is_alive():
|
||||
logger.debug("Reattaching to existing terminal session for instance %s", instance_id)
|
||||
return session
|
||||
else:
|
||||
# Session died, clean it up
|
||||
logger.debug("Existing session for instance %s is dead, cleaning up", instance_id)
|
||||
await session.close()
|
||||
del self._sessions[instance_id_str]
|
||||
|
||||
# Create new session
|
||||
logger.info("Creating new terminal session for instance %s", instance_id)
|
||||
session_id = str(uuid.uuid4())
|
||||
session = TerminalSession(session_id, instance_id, container_id, startup_command=startup_command)
|
||||
await session.start(startup_command=startup_command)
|
||||
self._sessions[instance_id_str] = session
|
||||
|
||||
return session
|
||||
|
||||
async def attach_websocket(
|
||||
self,
|
||||
session: TerminalSession,
|
||||
websocket: WebSocket,
|
||||
) -> None:
|
||||
"""Attach a WebSocket to an existing session."""
|
||||
# Handle concurrent connections - close existing ones
|
||||
if session.has_websockets():
|
||||
logger.debug("Closing existing WebSocket connections for instance %s", session.instance_id)
|
||||
for ws in list(session._websockets):
|
||||
try:
|
||||
await ws.close(code=4000, reason="New connection established")
|
||||
except Exception:
|
||||
pass
|
||||
session._websockets.clear()
|
||||
|
||||
# Attach new WebSocket
|
||||
session.attach_websocket(websocket)
|
||||
|
||||
# Replay buffer
|
||||
buffer = session.get_buffer()
|
||||
if buffer:
|
||||
try:
|
||||
await websocket.send_bytes(buffer)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def detach_websocket(
|
||||
self,
|
||||
session: TerminalSession,
|
||||
websocket: WebSocket,
|
||||
) -> None:
|
||||
"""Detach a WebSocket from a session."""
|
||||
session.detach_websocket(websocket)
|
||||
|
||||
async def reset_session(
|
||||
self,
|
||||
instance_id: uuid.UUID,
|
||||
container_id: str,
|
||||
startup_command: str | None = None,
|
||||
) -> TerminalSession:
|
||||
"""Reset a session by killing it and creating a new one."""
|
||||
instance_id_str = str(instance_id)
|
||||
|
||||
# Close existing session if any
|
||||
if instance_id_str in self._sessions:
|
||||
logger.debug("Resetting terminal session for instance %s", instance_id)
|
||||
old_session = self._sessions.pop(instance_id_str)
|
||||
await old_session.close()
|
||||
|
||||
# Create new session
|
||||
session_id = str(uuid.uuid4())
|
||||
session = TerminalSession(session_id, instance_id, container_id, startup_command=startup_command)
|
||||
await session.start(startup_command=startup_command)
|
||||
self._sessions[instance_id_str] = session
|
||||
|
||||
return session
|
||||
|
||||
async def close_all(self) -> None:
|
||||
"""Close all active sessions."""
|
||||
sessions = list(self._sessions.values())
|
||||
self._sessions.clear()
|
||||
for session in sessions:
|
||||
await session.close()
|
||||
|
||||
if self._idle_check_task and not self._idle_check_task.done():
|
||||
self._idle_check_task.cancel()
|
||||
|
||||
|
||||
# Global terminal manager instance
|
||||
terminal_manager = TerminalManager()
|
||||
@@ -0,0 +1,246 @@
|
||||
"""Terminal session management for tool instances."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import pty
|
||||
import select
|
||||
import signal
|
||||
import struct
|
||||
import fcntl
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TerminalSession:
|
||||
"""Manages a single terminal session connected to a docker container.
|
||||
|
||||
Supports persistent sessions that survive WebSocket disconnections.
|
||||
Multiple WebSocket connections can attach/detach from the same session.
|
||||
"""
|
||||
|
||||
# Circular buffer size (10KB)
|
||||
BUFFER_SIZE = 10 * 1024
|
||||
|
||||
# Idle timeout in seconds (30 minutes)
|
||||
IDLE_TIMEOUT = 30 * 60
|
||||
|
||||
def __init__(self, session_id: str, instance_id: uuid.UUID, container_id: str, startup_command: str | None = None) -> None:
|
||||
self.session_id = session_id
|
||||
self.instance_id = instance_id
|
||||
self.container_id = container_id
|
||||
self.startup_command = startup_command
|
||||
self.process: asyncio.subprocess.Process | None = None
|
||||
self._closed = False
|
||||
self._master_fd: int | None = None
|
||||
self._slave_fd: int | None = None
|
||||
|
||||
# Circular buffer for output replay
|
||||
self._output_buffer: deque[bytes] = deque(maxlen=self.BUFFER_SIZE)
|
||||
self._buffer_size = 0
|
||||
|
||||
# WebSocket connections
|
||||
self._websockets: set[Any] = set()
|
||||
|
||||
# Activity tracking
|
||||
self.last_activity = time.time()
|
||||
|
||||
# Terminal size
|
||||
self._cols = 80
|
||||
self._rows = 24
|
||||
|
||||
async def start(self, startup_command: str | None = None) -> None:
|
||||
"""Start the docker exec process with a shell using a PTY."""
|
||||
# Create a pseudo-terminal on the host
|
||||
self._master_fd, self._slave_fd = pty.openpty()
|
||||
|
||||
# Set the terminal size initially
|
||||
self._set_terminal_size(self._cols, self._rows)
|
||||
logger.debug(f"Starting terminal session {self.session_id} for container {self.container_id} with initial size {self._cols}x{self._rows}")
|
||||
|
||||
# Build the shell command
|
||||
if startup_command:
|
||||
shell_cmd = f'bash -c "{startup_command}" || true; exec bash -il'
|
||||
logger.debug(f"Using startup command for session {self.session_id}: {startup_command}")
|
||||
else:
|
||||
shell_cmd = "bash -il"
|
||||
|
||||
# Start docker exec with the slave fd as stdin/stdout/stderr
|
||||
# Using -it because the slave fd IS a TTY
|
||||
self.process = await asyncio.create_subprocess_exec(
|
||||
"docker",
|
||||
"exec",
|
||||
"-it",
|
||||
"-e",
|
||||
"TERM=xterm",
|
||||
self.container_id,
|
||||
"bash",
|
||||
"-c",
|
||||
shell_cmd,
|
||||
stdin=self._slave_fd,
|
||||
stdout=self._slave_fd,
|
||||
stderr=self._slave_fd,
|
||||
)
|
||||
|
||||
# Close slave fd in parent process
|
||||
os.close(self._slave_fd)
|
||||
self._slave_fd = None
|
||||
|
||||
self.last_activity = time.time()
|
||||
|
||||
def _set_terminal_size(self, cols: int, rows: int) -> None:
|
||||
"""Set the terminal size using TIOCSWINSZ."""
|
||||
if self._master_fd is None:
|
||||
logger.warning("Cannot resize: master_fd is None (session not started)")
|
||||
return
|
||||
# TIOCSWINSZ = 0x5414 on Linux
|
||||
TIOCSWINSZ = 0x5414
|
||||
size = struct.pack('HHHH', rows, cols, 0, 0)
|
||||
try:
|
||||
fcntl.ioctl(self._master_fd, TIOCSWINSZ, size)
|
||||
logger.debug(f"Resized PTY to {cols}x{rows} (fd={self._master_fd})")
|
||||
except (OSError, IOError) as e:
|
||||
logger.error(f"Failed to resize PTY: {e}")
|
||||
|
||||
async def read_output(self) -> bytes:
|
||||
"""Read output from the PTY master and store in buffer."""
|
||||
if self._master_fd is None or self._closed:
|
||||
return b""
|
||||
try:
|
||||
# Use select to check if data is available
|
||||
readable, _, _ = select.select([self._master_fd], [], [], 0.1)
|
||||
if readable:
|
||||
data = os.read(self._master_fd, 4096)
|
||||
if data:
|
||||
self._add_to_buffer(data)
|
||||
self.last_activity = time.time()
|
||||
return data
|
||||
return b""
|
||||
except (OSError, IOError, ValueError):
|
||||
return b""
|
||||
|
||||
def _add_to_buffer(self, data: bytes) -> None:
|
||||
"""Add data to circular buffer, maintaining size limit."""
|
||||
self._output_buffer.append(data)
|
||||
self._buffer_size += len(data)
|
||||
|
||||
# Trim if exceeds max size
|
||||
while self._buffer_size > self.BUFFER_SIZE and self._output_buffer:
|
||||
removed = self._output_buffer.popleft()
|
||||
self._buffer_size -= len(removed)
|
||||
|
||||
def get_buffer(self) -> bytes:
|
||||
"""Get buffered output for replay."""
|
||||
return b"".join(self._output_buffer)
|
||||
|
||||
async def write_input(self, data: bytes) -> None:
|
||||
"""Write input to the PTY master."""
|
||||
if self._master_fd is None or self._closed:
|
||||
return
|
||||
try:
|
||||
os.write(self._master_fd, data)
|
||||
self.last_activity = time.time()
|
||||
except (OSError, IOError):
|
||||
pass
|
||||
|
||||
async def resize(self, cols: int, rows: int) -> None:
|
||||
"""Resize the terminal."""
|
||||
if self._closed:
|
||||
logger.warning("Cannot resize: session is closed")
|
||||
return
|
||||
|
||||
# Only resize if dimensions actually changed
|
||||
if cols == self._cols and rows == self._rows:
|
||||
return
|
||||
|
||||
self._cols = cols
|
||||
self._rows = rows
|
||||
logger.debug(f"resize() called for session {self.session_id}: {cols}x{rows}")
|
||||
self._set_terminal_size(cols, rows)
|
||||
|
||||
# Docker exec -it creates its own PTY inside the container,
|
||||
# so host PTY resize doesn't propagate to the container shell.
|
||||
# Send SIGWINCH to the docker exec process on the host.
|
||||
# Docker exec forwards signals to the container process, which should
|
||||
# cause the container's shell to re-read its terminal size.
|
||||
if self.process and self.process.pid:
|
||||
try:
|
||||
os.kill(self.process.pid, signal.SIGWINCH)
|
||||
logger.debug(f"Sent SIGWINCH to docker exec process {self.process.pid} for session {self.session_id}")
|
||||
except ProcessLookupError:
|
||||
logger.warning(f"docker exec process {self.process.pid} not found for session {self.session_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to send SIGWINCH: {e}")
|
||||
|
||||
async def reset(self) -> None:
|
||||
"""Reset the session by killing the process and clearing state."""
|
||||
await self.close()
|
||||
self._closed = False
|
||||
self._output_buffer.clear()
|
||||
self._buffer_size = 0
|
||||
self._websockets.clear()
|
||||
self.process = None
|
||||
self._master_fd = None
|
||||
self._slave_fd = None
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the session and cleanup."""
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
|
||||
if self._master_fd is not None:
|
||||
try:
|
||||
os.close(self._master_fd)
|
||||
except OSError:
|
||||
pass
|
||||
self._master_fd = None
|
||||
|
||||
if self.process is not None:
|
||||
try:
|
||||
self.process.kill()
|
||||
await asyncio.wait_for(self.process.wait(), timeout=2.0)
|
||||
except (asyncio.TimeoutError, ProcessLookupError):
|
||||
pass
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
"""Check if the session process is still running."""
|
||||
if self.process is None:
|
||||
return False
|
||||
return self.process.returncode is None
|
||||
|
||||
def is_idle(self) -> bool:
|
||||
"""Check if the session has been idle for too long."""
|
||||
if self._websockets:
|
||||
return False
|
||||
return time.time() - self.last_activity > self.IDLE_TIMEOUT
|
||||
|
||||
def attach_websocket(self, websocket: Any) -> None:
|
||||
"""Attach a WebSocket to this session."""
|
||||
self._websockets.add(websocket)
|
||||
self.last_activity = time.time()
|
||||
|
||||
def detach_websocket(self, websocket: Any) -> None:
|
||||
"""Detach a WebSocket from this session."""
|
||||
self._websockets.discard(websocket)
|
||||
|
||||
def has_websockets(self) -> bool:
|
||||
"""Check if any WebSockets are attached."""
|
||||
return len(self._websockets) > 0
|
||||
|
||||
async def send_to_all(self, data: bytes) -> None:
|
||||
"""Send data to all attached WebSockets."""
|
||||
dead_sockets = set()
|
||||
for ws in self._websockets:
|
||||
try:
|
||||
await ws.send_bytes(data)
|
||||
except Exception:
|
||||
dead_sockets.add(ws)
|
||||
|
||||
# Clean up dead sockets
|
||||
for ws in dead_sockets:
|
||||
self._websockets.discard(ws)
|
||||
@@ -0,0 +1,313 @@
|
||||
"""Git control utilities for repository operations."""
|
||||
|
||||
import subprocess
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
def _run_git_command(repo_path: str, *args: str) -> str:
|
||||
"""Run a git command in the repository directory."""
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=repo_path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Git command failed: {result.stderr}")
|
||||
return result.stdout
|
||||
|
||||
|
||||
@dataclass
|
||||
class GitStatus:
|
||||
"""Represents the working directory status."""
|
||||
|
||||
branch: str
|
||||
modified: list[str] = field(default_factory=list)
|
||||
added: list[str] = field(default_factory=list)
|
||||
deleted: list[str] = field(default_factory=list)
|
||||
untracked: list[str] = field(default_factory=list)
|
||||
renamed: list[str] = field(default_factory=list)
|
||||
ahead: int = 0
|
||||
behind: int = 0
|
||||
|
||||
|
||||
def get_status(repo_path: str) -> GitStatus:
|
||||
"""Get the working directory status.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
|
||||
Returns:
|
||||
GitStatus with changes
|
||||
"""
|
||||
# Get current branch
|
||||
try:
|
||||
branch = _run_git_command(repo_path, "rev-parse", "--abbrev-ref", "HEAD").strip()
|
||||
except RuntimeError:
|
||||
try:
|
||||
branch = _run_git_command(repo_path, "symbolic-ref", "--short", "HEAD").strip()
|
||||
except RuntimeError:
|
||||
branch = "HEAD"
|
||||
|
||||
status = GitStatus(branch=branch)
|
||||
|
||||
# Get status with porcelain format
|
||||
try:
|
||||
output = _run_git_command(repo_path, "status", "--porcelain", "--branch")
|
||||
except RuntimeError:
|
||||
return status
|
||||
|
||||
for line in output.strip().split("\n"):
|
||||
if not line:
|
||||
continue
|
||||
|
||||
# Branch info line starts with ##
|
||||
if line.startswith("## "):
|
||||
branch_info = line[3:]
|
||||
# Parse ahead/behind info
|
||||
if "[ahead " in branch_info:
|
||||
ahead_str = branch_info.split("[ahead ")[1].split("]")[0]
|
||||
status.ahead = int(ahead_str.split(",")[0])
|
||||
if "[behind " in branch_info:
|
||||
behind_str = branch_info.split("[behind ")[1].split("]")[0]
|
||||
status.behind = int(behind_str.split(",")[0])
|
||||
continue
|
||||
|
||||
# Parse status code
|
||||
if len(line) < 3:
|
||||
continue
|
||||
|
||||
index_status = line[0]
|
||||
worktree_status = line[1]
|
||||
filename = line[3:]
|
||||
|
||||
# Untracked files
|
||||
if index_status == "?" and worktree_status == "?":
|
||||
status.untracked.append(filename)
|
||||
continue
|
||||
|
||||
# Added files
|
||||
if index_status == "A" or worktree_status == "A":
|
||||
status.added.append(filename)
|
||||
continue
|
||||
|
||||
# Deleted files
|
||||
if index_status == "D" or worktree_status == "D":
|
||||
status.deleted.append(filename)
|
||||
continue
|
||||
|
||||
# Renamed files
|
||||
if index_status == "R" or worktree_status == "R":
|
||||
status.renamed.append(filename)
|
||||
continue
|
||||
|
||||
# Modified files
|
||||
if index_status == "M" or worktree_status == "M":
|
||||
status.modified.append(filename)
|
||||
continue
|
||||
|
||||
return status
|
||||
|
||||
|
||||
def create_branch(repo_path: str, name: str, base_branch: str = "HEAD") -> None:
|
||||
"""Create a new branch.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
name: Branch name
|
||||
base_branch: Base branch to create from (default: HEAD)
|
||||
|
||||
Raises:
|
||||
RuntimeError: If branch creation fails
|
||||
"""
|
||||
try:
|
||||
_run_git_command(repo_path, "rev-parse", "--verify", "HEAD^{commit}")
|
||||
except RuntimeError:
|
||||
# No commits yet - empty repository
|
||||
try:
|
||||
_run_git_command(repo_path, "checkout", "--orphan", name)
|
||||
except RuntimeError as e:
|
||||
if "work tree" in str(e).lower():
|
||||
# Bare repository - use symbolic-ref instead
|
||||
_run_git_command(repo_path, "symbolic-ref", "HEAD", f"refs/heads/{name}")
|
||||
return
|
||||
raise
|
||||
return
|
||||
|
||||
_run_git_command(repo_path, "branch", name, base_branch)
|
||||
|
||||
|
||||
def delete_branch(repo_path: str, name: str, force: bool = False) -> None:
|
||||
"""Delete a branch.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
name: Branch name
|
||||
force: Force delete even if not merged
|
||||
|
||||
Raises:
|
||||
RuntimeError: If branch deletion fails
|
||||
"""
|
||||
flag = "-D" if force else "-d"
|
||||
_run_git_command(repo_path, "branch", flag, name)
|
||||
|
||||
|
||||
def checkout_branch(repo_path: str, name: str) -> None:
|
||||
"""Checkout a branch.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
name: Branch name
|
||||
|
||||
Raises:
|
||||
RuntimeError: If checkout fails
|
||||
"""
|
||||
try:
|
||||
_run_git_command(repo_path, "checkout", name)
|
||||
except RuntimeError as e:
|
||||
if "work tree" in str(e).lower():
|
||||
# Bare repository - use symbolic-ref instead
|
||||
_run_git_command(repo_path, "symbolic-ref", "HEAD", f"refs/heads/{name}")
|
||||
return
|
||||
raise
|
||||
|
||||
|
||||
def commit_changes(
|
||||
repo_path: str,
|
||||
message: str,
|
||||
author_name: str,
|
||||
author_email: str,
|
||||
files: list[str] | None = None,
|
||||
) -> str:
|
||||
"""Commit changes to the repository.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
message: Commit message
|
||||
author_name: Author name
|
||||
author_email: Author email
|
||||
files: Specific files to commit (None = all staged)
|
||||
|
||||
Returns:
|
||||
Commit hash
|
||||
|
||||
Raises:
|
||||
RuntimeError: If commit fails
|
||||
"""
|
||||
# Stage files if specified
|
||||
if files:
|
||||
for file in files:
|
||||
_run_git_command(repo_path, "add", file)
|
||||
else:
|
||||
_run_git_command(repo_path, "add", "-A")
|
||||
|
||||
# Commit
|
||||
_run_git_command(
|
||||
repo_path,
|
||||
"commit",
|
||||
"-m",
|
||||
message,
|
||||
f"--author={author_name} <{author_email}>",
|
||||
)
|
||||
|
||||
# Return commit hash
|
||||
return _run_git_command(repo_path, "rev-parse", "HEAD").strip()
|
||||
|
||||
|
||||
def fetch(repo_path: str) -> None:
|
||||
"""Fetch from remote.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
|
||||
Raises:
|
||||
RuntimeError: If fetch fails
|
||||
"""
|
||||
_run_git_command(repo_path, "fetch", "--all")
|
||||
|
||||
|
||||
def pull(repo_path: str, branch: str | None = None) -> None:
|
||||
"""Pull updates from remote.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
branch: Branch to pull (default: current branch)
|
||||
|
||||
Raises:
|
||||
RuntimeError: If pull fails
|
||||
"""
|
||||
args = ["pull"]
|
||||
if branch:
|
||||
args.append("origin")
|
||||
args.append(branch)
|
||||
_run_git_command(repo_path, *args)
|
||||
|
||||
|
||||
def push(repo_path: str, branch: str | None = None) -> None:
|
||||
"""Push changes to remote.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
branch: Branch to push (default: current branch)
|
||||
|
||||
Raises:
|
||||
RuntimeError: If push fails
|
||||
"""
|
||||
args = ["push"]
|
||||
if branch:
|
||||
args.extend(["origin", branch])
|
||||
_run_git_command(repo_path, *args)
|
||||
|
||||
|
||||
def merge(
|
||||
repo_path: str,
|
||||
source_branch: str,
|
||||
target_branch: str | None = None,
|
||||
message: str | None = None,
|
||||
) -> str:
|
||||
"""Merge a branch into the current branch.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
source_branch: Branch to merge from
|
||||
target_branch: Branch to merge into (default: current branch)
|
||||
message: Merge commit message
|
||||
|
||||
Returns:
|
||||
Merge commit hash
|
||||
|
||||
Raises:
|
||||
RuntimeError: If merge fails (including conflicts)
|
||||
"""
|
||||
# Checkout target branch if specified
|
||||
if target_branch:
|
||||
checkout_branch(repo_path, target_branch)
|
||||
|
||||
# Merge
|
||||
args = ["merge", source_branch]
|
||||
if message:
|
||||
args.extend(["-m", message])
|
||||
|
||||
_run_git_command(repo_path, *args)
|
||||
|
||||
# Return merge commit hash
|
||||
return _run_git_command(repo_path, "rev-parse", "HEAD").strip()
|
||||
|
||||
|
||||
def get_current_branch(repo_path: str) -> str:
|
||||
"""Get the current branch name.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
|
||||
Returns:
|
||||
Current branch name
|
||||
"""
|
||||
try:
|
||||
branch = _run_git_command(repo_path, "rev-parse", "--abbrev-ref", "HEAD").strip()
|
||||
if branch != "HEAD":
|
||||
return branch
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
return _run_git_command(repo_path, "symbolic-ref", "--short", "HEAD").strip()
|
||||
@@ -0,0 +1,457 @@
|
||||
"""Git file utilities for browsing repository contents."""
|
||||
|
||||
import logging
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class FileTreeEntry:
|
||||
"""Represents a file or directory in the repository."""
|
||||
|
||||
name: str
|
||||
type: str # "file" or "directory"
|
||||
path: str
|
||||
size: int | None = None
|
||||
mode: str | None = None
|
||||
last_commit: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class BranchInfo:
|
||||
"""Represents a git branch."""
|
||||
|
||||
name: str
|
||||
is_default: bool
|
||||
last_commit: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class FileContent:
|
||||
"""Represents file content and metadata."""
|
||||
|
||||
path: str
|
||||
branch: str
|
||||
content: str
|
||||
size: int
|
||||
encoding: str
|
||||
language: str | None
|
||||
is_binary: bool
|
||||
last_commit: dict[str, Any] | None = None
|
||||
|
||||
|
||||
def _run_git_command(repo_path: str, *args: str) -> str:
|
||||
"""Run a git command in the repository directory."""
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=repo_path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
stderr = result.stderr
|
||||
# Handle "dubious ownership" security error
|
||||
if "dubious ownership" in stderr.lower():
|
||||
logger.warning("Git ownership mismatch for %s, adding to safe.directory", repo_path)
|
||||
# Add this directory to git's safe.directory list
|
||||
subprocess.run(
|
||||
["git", "config", "--global", "--add", "safe.directory", repo_path],
|
||||
capture_output=True,
|
||||
)
|
||||
# Retry the command
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=repo_path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
return result.stdout
|
||||
stderr = result.stderr
|
||||
raise RuntimeError(f"Git command failed: {stderr}")
|
||||
return result.stdout
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def list_tree(repo_path: str, branch: str = "main", path: str = "") -> list[FileTreeEntry]:
|
||||
"""List files and directories in a repository path.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
branch: Branch name to list from
|
||||
path: Directory path within the repository (empty for root)
|
||||
|
||||
Returns:
|
||||
List of FileTreeEntry objects
|
||||
"""
|
||||
tree_path = f"{branch}:{path}" if path else branch
|
||||
|
||||
try:
|
||||
output = _run_git_command(repo_path, "ls-tree", "-l", tree_path)
|
||||
except RuntimeError as e:
|
||||
logger.warning("git ls-tree failed for %s with branch '%s': %s", repo_path, tree_path, str(e))
|
||||
# Try with HEAD if branch doesn't exist
|
||||
tree_path = f"HEAD:{path}" if path else "HEAD"
|
||||
try:
|
||||
output = _run_git_command(repo_path, "ls-tree", "-l", tree_path)
|
||||
except RuntimeError as e:
|
||||
logger.error("git ls-tree failed for %s with HEAD: %s", repo_path, str(e))
|
||||
# Check if this is an empty repository (no commits yet)
|
||||
error_msg = str(e).lower()
|
||||
if "not a valid object name" in error_msg or "does not exist" in error_msg:
|
||||
# Empty repository - return empty list
|
||||
return []
|
||||
raise
|
||||
|
||||
entries = []
|
||||
for line in output.strip().split("\n"):
|
||||
if not line:
|
||||
continue
|
||||
|
||||
# Format: <mode> <type> <hash> <size>\t<name>
|
||||
parts = line.split("\t", 1)
|
||||
if len(parts) != 2:
|
||||
continue
|
||||
|
||||
meta, name = parts
|
||||
meta_parts = meta.split()
|
||||
if len(meta_parts) < 4:
|
||||
continue
|
||||
|
||||
mode = meta_parts[0]
|
||||
obj_type = meta_parts[1]
|
||||
_ = meta_parts[2] # object hash, not used
|
||||
size = int(meta_parts[3]) if obj_type == "blob" else None
|
||||
|
||||
entry_path = f"{path}/{name}" if path else name
|
||||
|
||||
# Get last commit info for this entry
|
||||
last_commit = _get_last_commit_for_path(repo_path, branch, entry_path)
|
||||
|
||||
entries.append(
|
||||
FileTreeEntry(
|
||||
name=name,
|
||||
type="directory" if obj_type == "tree" else "file",
|
||||
path=entry_path,
|
||||
size=size,
|
||||
mode=mode,
|
||||
last_commit=last_commit,
|
||||
)
|
||||
)
|
||||
|
||||
return entries
|
||||
|
||||
|
||||
def _get_last_commit_for_path(repo_path: str, branch: str, path: str) -> dict[str, Any] | None:
|
||||
"""Get the last commit that modified a path."""
|
||||
try:
|
||||
output = _run_git_command(
|
||||
repo_path,
|
||||
"log",
|
||||
"-1",
|
||||
"--format=%H|%s|%an|%aI",
|
||||
branch,
|
||||
"--",
|
||||
path,
|
||||
)
|
||||
if not output.strip():
|
||||
return None
|
||||
|
||||
parts = output.strip().split("|", 3)
|
||||
if len(parts) != 4:
|
||||
return None
|
||||
|
||||
return {
|
||||
"hash": parts[0],
|
||||
"message": parts[1],
|
||||
"author": parts[2],
|
||||
"date": parts[3],
|
||||
}
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
|
||||
def get_file_content(repo_path: str, branch: str, path: str) -> FileContent:
|
||||
"""Get the content of a file.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
branch: Branch name
|
||||
path: File path within the repository
|
||||
|
||||
Returns:
|
||||
FileContent with content and metadata
|
||||
"""
|
||||
# Check if file exists
|
||||
try:
|
||||
_run_git_command(repo_path, "cat-file", "-e", f"{branch}:{path}")
|
||||
except RuntimeError:
|
||||
raise FileNotFoundError(f"File '{path}' not found in branch '{branch}'")
|
||||
|
||||
# Get file size
|
||||
size_output = _run_git_command(repo_path, "cat-file", "-s", f"{branch}:{path}")
|
||||
size = int(size_output.strip())
|
||||
|
||||
# Check if binary
|
||||
is_binary = _is_binary_file(repo_path, branch, path)
|
||||
|
||||
# Get content (only for text files)
|
||||
content = ""
|
||||
if not is_binary:
|
||||
content = _run_git_command(repo_path, "show", f"{branch}:{path}")
|
||||
|
||||
# Detect language from extension
|
||||
language = _detect_language(path)
|
||||
|
||||
# Get last commit
|
||||
last_commit = _get_last_commit_for_path(repo_path, branch, path)
|
||||
|
||||
return FileContent(
|
||||
path=path,
|
||||
branch=branch,
|
||||
content=content,
|
||||
size=size,
|
||||
encoding="utf-8",
|
||||
language=language,
|
||||
is_binary=is_binary,
|
||||
last_commit=last_commit,
|
||||
)
|
||||
|
||||
|
||||
def _is_binary_file(repo_path: str, branch: str, path: str) -> bool:
|
||||
"""Check if a file is binary using raw bytes to avoid encoding issues."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "show", f"{branch}:{path}"],
|
||||
cwd=repo_path,
|
||||
capture_output=True,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Git command failed: {result.stderr.decode()}")
|
||||
# A file is binary if it contains null bytes
|
||||
return b"\x00" in result.stdout
|
||||
except RuntimeError:
|
||||
return True
|
||||
|
||||
|
||||
def _detect_language(path: str) -> str | None:
|
||||
"""Detect programming language from file extension."""
|
||||
ext = Path(path).suffix.lower()
|
||||
language_map = {
|
||||
".py": "python",
|
||||
".js": "javascript",
|
||||
".ts": "typescript",
|
||||
".jsx": "jsx",
|
||||
".tsx": "tsx",
|
||||
".html": "html",
|
||||
".css": "css",
|
||||
".scss": "scss",
|
||||
".json": "json",
|
||||
".md": "markdown",
|
||||
".yaml": "yaml",
|
||||
".yml": "yaml",
|
||||
".sh": "bash",
|
||||
".rs": "rust",
|
||||
".go": "go",
|
||||
".java": "java",
|
||||
".c": "c",
|
||||
".cpp": "cpp",
|
||||
".h": "c",
|
||||
".php": "php",
|
||||
".rb": "ruby",
|
||||
".sql": "sql",
|
||||
".dockerfile": "dockerfile",
|
||||
".vue": "vue",
|
||||
".svelte": "svelte",
|
||||
}
|
||||
return language_map.get(ext)
|
||||
|
||||
|
||||
def list_branches(repo_path: str) -> tuple[list[BranchInfo], str]:
|
||||
"""List all branches and identify the default branch.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
|
||||
Returns:
|
||||
Tuple of (list of BranchInfo, default branch name)
|
||||
"""
|
||||
# Get all branches
|
||||
try:
|
||||
output = _run_git_command(repo_path, "branch", "-a", "--format=%(refname:short)")
|
||||
except RuntimeError as e:
|
||||
logger.error("Failed to list branches for %s: %s", repo_path, str(e))
|
||||
raise
|
||||
|
||||
branches: list[BranchInfo] = []
|
||||
default_branch = "main"
|
||||
|
||||
# Get list of remote names to properly filter remote tracking branches
|
||||
try:
|
||||
remote_output = _run_git_command(repo_path, "remote")
|
||||
remote_names = {r.strip() for r in remote_output.strip().split("\n") if r.strip()}
|
||||
except RuntimeError:
|
||||
remote_names = set()
|
||||
|
||||
for line in output.strip().split("\n"):
|
||||
if not line:
|
||||
continue
|
||||
|
||||
branch_name = line.strip()
|
||||
|
||||
# Skip detached HEAD pointer
|
||||
if branch_name == "HEAD":
|
||||
continue
|
||||
|
||||
# Skip remote tracking branches - they appear as "origin/branch-name"
|
||||
# Check if first part is a remote name
|
||||
if "/" in branch_name:
|
||||
first_part = branch_name.split("/", 1)[0]
|
||||
if first_part in remote_names:
|
||||
# Extract just the branch name part (after "origin/")
|
||||
branch_name = branch_name.split("/", 1)[1]
|
||||
elif branch_name.startswith("remotes/"):
|
||||
# Handle "remotes/origin/branch-name" format
|
||||
parts = branch_name.split("/", 2)
|
||||
if len(parts) >= 3:
|
||||
branch_name = parts[2]
|
||||
else:
|
||||
continue
|
||||
|
||||
# Skip duplicates
|
||||
if any(b.name == branch_name for b in branches):
|
||||
continue
|
||||
|
||||
# Check if this is the default branch (HEAD points to it)
|
||||
try:
|
||||
head_output = _run_git_command(
|
||||
repo_path,
|
||||
"symbolic-ref",
|
||||
"HEAD",
|
||||
)
|
||||
if head_output.strip() == f"refs/heads/{branch_name}":
|
||||
default_branch = branch_name
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
# Get last commit for branch
|
||||
last_commit = _get_last_commit_for_path(repo_path, branch_name, ".")
|
||||
|
||||
branches.append(
|
||||
BranchInfo(
|
||||
name=branch_name,
|
||||
is_default=(branch_name == default_branch),
|
||||
last_commit=last_commit,
|
||||
)
|
||||
)
|
||||
|
||||
# If no branches found, try to get HEAD
|
||||
if not branches:
|
||||
try:
|
||||
output = _run_git_command(repo_path, "rev-parse", "--abbrev-ref", "HEAD")
|
||||
branch_name = output.strip()
|
||||
if branch_name and branch_name != "HEAD":
|
||||
last_commit = _get_last_commit_for_path(repo_path, branch_name, ".")
|
||||
branches.append(
|
||||
BranchInfo(
|
||||
name=branch_name,
|
||||
is_default=True,
|
||||
last_commit=last_commit,
|
||||
)
|
||||
)
|
||||
default_branch = branch_name
|
||||
except RuntimeError:
|
||||
try:
|
||||
output = _run_git_command(repo_path, "symbolic-ref", "--short", "HEAD")
|
||||
branch_name = output.strip()
|
||||
if branch_name:
|
||||
branches.append(
|
||||
BranchInfo(
|
||||
name=branch_name,
|
||||
is_default=True,
|
||||
last_commit=None,
|
||||
)
|
||||
)
|
||||
default_branch = branch_name
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
return branches, default_branch
|
||||
|
||||
|
||||
def commit_file(
|
||||
repo_path: str,
|
||||
branch: str,
|
||||
path: str,
|
||||
content: str,
|
||||
commit_message: str,
|
||||
author_name: str,
|
||||
author_email: str,
|
||||
) -> str:
|
||||
"""Commit a file change.
|
||||
|
||||
Args:
|
||||
repo_path: Path to the git repository
|
||||
branch: Branch to commit to
|
||||
path: File path within the repository
|
||||
content: New file content
|
||||
commit_message: Commit message
|
||||
author_name: Author name
|
||||
author_email: Author email
|
||||
|
||||
Returns:
|
||||
Commit hash
|
||||
"""
|
||||
# For bare repositories, we need to use git commands differently
|
||||
# We'll create a temporary worktree, make changes, and commit
|
||||
|
||||
import tempfile
|
||||
import os
|
||||
|
||||
# Create a temporary worktree
|
||||
with tempfile.TemporaryDirectory() as worktree_path:
|
||||
# Add worktree
|
||||
_run_git_command(
|
||||
repo_path,
|
||||
"worktree",
|
||||
"add",
|
||||
"--detach",
|
||||
worktree_path,
|
||||
branch,
|
||||
)
|
||||
|
||||
try:
|
||||
# Write file content
|
||||
file_path = os.path.join(worktree_path, path)
|
||||
os.makedirs(os.path.dirname(file_path), exist_ok=True)
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
|
||||
# Configure git author
|
||||
_run_git_command(worktree_path, "config", "user.name", author_name)
|
||||
_run_git_command(worktree_path, "config", "user.email", author_email)
|
||||
|
||||
# Stage and commit
|
||||
_run_git_command(worktree_path, "add", path)
|
||||
_run_git_command(
|
||||
worktree_path,
|
||||
"commit",
|
||||
"-m",
|
||||
commit_message,
|
||||
)
|
||||
|
||||
# Get commit hash
|
||||
commit_hash = _run_git_command(
|
||||
worktree_path,
|
||||
"rev-parse",
|
||||
"HEAD",
|
||||
).strip()
|
||||
|
||||
return commit_hash
|
||||
|
||||
finally:
|
||||
# Remove worktree
|
||||
_run_git_command(repo_path, "worktree", "remove", worktree_path)
|
||||
@@ -0,0 +1,382 @@
|
||||
"""Git history extraction utilities for bare/mirror repositories."""
|
||||
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class Commit:
|
||||
"""Represents a single git commit."""
|
||||
|
||||
hash: str
|
||||
short_hash: str
|
||||
parents: list[str]
|
||||
author: str
|
||||
email: str
|
||||
date: str
|
||||
timestamp: int
|
||||
message: str
|
||||
branches: list[str]
|
||||
tags: list[str]
|
||||
|
||||
|
||||
@dataclass
|
||||
class FileChange:
|
||||
"""Represents a changed file in a commit."""
|
||||
|
||||
path: str
|
||||
change_type: str
|
||||
insertions: int
|
||||
deletions: int
|
||||
diff: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class CommitDetail(Commit):
|
||||
"""Extended commit info with diff."""
|
||||
|
||||
body: str
|
||||
stats: dict[str, int]
|
||||
files: list[FileChange]
|
||||
|
||||
|
||||
def _run_git_command(repo_path: str, args: list[str]) -> str:
|
||||
"""Execute a git command in the repository directory."""
|
||||
result = subprocess.run(
|
||||
["git", "-C", repo_path, *args],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Git command failed: {result.stderr}")
|
||||
return result.stdout
|
||||
|
||||
|
||||
def get_commit_history(repo_path: str, branch: str | None = None, limit: int = 100, offset: int = 0) -> dict[str, Any]:
|
||||
"""Extract commit history from a git repository.
|
||||
|
||||
Returns structured data including commits, branches, and graph information.
|
||||
"""
|
||||
# Get list of branches
|
||||
branches_output = _run_git_command(repo_path, ["branch", "-a", "--format=%(refname:short)"])
|
||||
branches = [b.strip() for b in branches_output.strip().split("\n") if b.strip()]
|
||||
|
||||
# Build git log command - use NULL bytes as separators to avoid parsing issues
|
||||
log_args = [
|
||||
"log",
|
||||
"--format=%H%x00%P%x00%an%x00%ae%x00%at%x00%s",
|
||||
f"--max-count={limit}",
|
||||
f"--skip={offset}",
|
||||
]
|
||||
if branch:
|
||||
log_args.append(branch)
|
||||
else:
|
||||
log_args.append("--all")
|
||||
|
||||
log_output = _run_git_command(repo_path, log_args)
|
||||
|
||||
# Get branch info for each commit
|
||||
branch_map = _get_branch_map(repo_path)
|
||||
tag_map = _get_tag_map(repo_path)
|
||||
|
||||
commits = []
|
||||
log_lines = log_output.strip().split("\n") if log_output.strip() else []
|
||||
|
||||
for line in log_lines:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
|
||||
parts = line.split("\x00")
|
||||
if len(parts) < 6:
|
||||
continue
|
||||
|
||||
commit_hash = parts[0]
|
||||
parents = parts[1].split() if parts[1] else []
|
||||
|
||||
commits.append(
|
||||
Commit(
|
||||
hash=commit_hash,
|
||||
short_hash=commit_hash[:7],
|
||||
parents=parents,
|
||||
author=parts[2],
|
||||
email=parts[3],
|
||||
date=parts[4],
|
||||
timestamp=int(parts[4]),
|
||||
message=parts[5],
|
||||
branches=branch_map.get(commit_hash, []),
|
||||
tags=tag_map.get(commit_hash, []),
|
||||
)
|
||||
)
|
||||
|
||||
# Get total commit count
|
||||
count_output = _run_git_command(repo_path, ["rev-list", "--all", "--count"])
|
||||
total_commits = int(count_output.strip()) if count_output.strip() else 0
|
||||
|
||||
# Build graph data and generate graph symbols
|
||||
graph_data = _build_graph_data(commits)
|
||||
|
||||
# Generate simple graph symbols based on parent count
|
||||
commit_dicts = []
|
||||
for i, commit in enumerate(commits):
|
||||
if len(commit.parents) == 0:
|
||||
graph_symbol = "○" # Initial commit
|
||||
elif len(commit.parents) > 1:
|
||||
graph_symbol = "●" # Merge commit
|
||||
else:
|
||||
graph_symbol = "○" # Regular commit
|
||||
|
||||
# Simple depth calculation based on merge status
|
||||
graph_depth = min(len(commit.parents), 3)
|
||||
|
||||
commit_dicts.append(_commit_to_dict(commit, graph_symbol, graph_depth))
|
||||
|
||||
return {
|
||||
"commits": commit_dicts,
|
||||
"branches": branches,
|
||||
"total_commits": total_commits,
|
||||
"graph_data": graph_data,
|
||||
}
|
||||
|
||||
|
||||
def get_commit_detail(repo_path: str, commit_hash: str) -> dict[str, Any]:
|
||||
"""Get detailed information about a specific commit."""
|
||||
# Get commit metadata
|
||||
format_str = "%H|%P|%an|%ae|%at|%s|%b"
|
||||
log_output = _run_git_command(
|
||||
repo_path, ["log", "-1", f"--format={format_str}", commit_hash]
|
||||
)
|
||||
|
||||
parts = log_output.strip().split("|", 6)
|
||||
if len(parts) < 6:
|
||||
raise ValueError(f"Invalid commit: {commit_hash}")
|
||||
|
||||
commit_hash = parts[0]
|
||||
parents = parts[1].split() if parts[1] else []
|
||||
author = parts[2]
|
||||
email = parts[3]
|
||||
timestamp = int(parts[4])
|
||||
message = parts[5]
|
||||
body = parts[6] if len(parts) > 6 else ""
|
||||
|
||||
# Get stats
|
||||
stat_output = _run_git_command(
|
||||
repo_path, ["show", "--stat", "--format=", commit_hash]
|
||||
)
|
||||
stats = _parse_stats(stat_output)
|
||||
|
||||
# Get diff
|
||||
diff_output = _run_git_command(
|
||||
repo_path, ["show", "--format=", commit_hash]
|
||||
)
|
||||
files = _parse_diff(diff_output)
|
||||
|
||||
# Get branch/tag info
|
||||
branch_map = _get_branch_map(repo_path)
|
||||
tag_map = _get_tag_map(repo_path)
|
||||
|
||||
return {
|
||||
"hash": commit_hash,
|
||||
"short_hash": commit_hash[:7],
|
||||
"parents": parents,
|
||||
"author_name": author,
|
||||
"author_email": email,
|
||||
"author_date": datetime.fromtimestamp(timestamp, tz=timezone.utc).isoformat(),
|
||||
"committer_name": author, # TODO: extract committer separately
|
||||
"committer_email": email, # TODO: extract committer separately
|
||||
"committer_date": datetime.fromtimestamp(timestamp, tz=timezone.utc).isoformat(),
|
||||
"message": message,
|
||||
"body": body,
|
||||
"branches": branch_map.get(commit_hash, []),
|
||||
"tags": tag_map.get(commit_hash, []),
|
||||
"stats": stats,
|
||||
"diff": diff_output,
|
||||
"files": [_file_change_to_dict(f) for f in files],
|
||||
}
|
||||
|
||||
|
||||
def _get_branch_map(repo_path: str) -> dict[str, list[str]]:
|
||||
"""Build a mapping of commit hash to branch names."""
|
||||
result = {}
|
||||
branch_output = _run_git_command(
|
||||
repo_path, ["for-each-ref", "--format=%(objectname) %(refname:short)", "refs/heads/"]
|
||||
)
|
||||
|
||||
for line in branch_output.strip().split("\n"):
|
||||
if " " in line:
|
||||
commit_hash, branch_name = line.split(" ", 1)
|
||||
if commit_hash not in result:
|
||||
result[commit_hash] = []
|
||||
result[commit_hash].append(branch_name)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _get_tag_map(repo_path: str) -> dict[str, list[str]]:
|
||||
"""Build a mapping of commit hash to tag names."""
|
||||
result = {}
|
||||
tag_output = _run_git_command(
|
||||
repo_path, ["for-each-ref", "--format=%(objectname) %(refname:short)", "refs/tags/"]
|
||||
)
|
||||
|
||||
for line in tag_output.strip().split("\n"):
|
||||
if " " in line:
|
||||
commit_hash, tag_name = line.split(" ", 1)
|
||||
if commit_hash not in result:
|
||||
result[commit_hash] = []
|
||||
result[commit_hash].append(tag_name)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _build_graph_data(commits: list[Commit]) -> dict[str, Any]:
|
||||
"""Build graph visualization data from commits."""
|
||||
if not commits:
|
||||
return {"nodes": [], "edges": []}
|
||||
|
||||
# Create hash to index mapping
|
||||
hash_to_idx = {c.hash: i for i, c in enumerate(commits)}
|
||||
|
||||
nodes = []
|
||||
edges = []
|
||||
|
||||
for i, commit in enumerate(commits):
|
||||
# Calculate column based on branch
|
||||
column = 0
|
||||
if commit.branches:
|
||||
# Use first branch as column indicator
|
||||
column = hash(commit.branches[0]) % 5
|
||||
|
||||
nodes.append(
|
||||
{
|
||||
"hash": commit.hash,
|
||||
"x": column * 60 + 30,
|
||||
"y": i * 50 + 25,
|
||||
"column": column,
|
||||
}
|
||||
)
|
||||
|
||||
# Create edges to parents
|
||||
for parent_hash in commit.parents:
|
||||
if parent_hash in hash_to_idx:
|
||||
edges.append(
|
||||
{
|
||||
"from_hash": commit.hash,
|
||||
"to_hash": parent_hash,
|
||||
"type": "parent",
|
||||
}
|
||||
)
|
||||
|
||||
return {"nodes": nodes, "edges": edges}
|
||||
|
||||
|
||||
def _parse_stats(stat_output: str) -> dict[str, int]:
|
||||
"""Parse git show --stat output."""
|
||||
lines = stat_output.strip().split("\n")
|
||||
stats = {"files_changed": 0, "insertions": 0, "deletions": 0}
|
||||
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
if "files changed" in line or "file changed" in line:
|
||||
# Parse summary line like "3 files changed, 45 insertions(+), 12 deletions(-)"
|
||||
parts = line.split(",")
|
||||
for part in parts:
|
||||
part = part.strip()
|
||||
if "file" in part:
|
||||
try:
|
||||
stats["files_changed"] = int(part.split()[0])
|
||||
except (ValueError, IndexError):
|
||||
pass
|
||||
elif "insertion" in part:
|
||||
try:
|
||||
stats["insertions"] = int(part.split()[0])
|
||||
except (ValueError, IndexError):
|
||||
pass
|
||||
elif "deletion" in part:
|
||||
try:
|
||||
stats["deletions"] = int(part.split()[0])
|
||||
except (ValueError, IndexError):
|
||||
pass
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
def _parse_diff(diff_output: str) -> list[FileChange]:
|
||||
"""Parse git diff output into file changes."""
|
||||
files = []
|
||||
current_file = None
|
||||
current_diff = []
|
||||
|
||||
for line in diff_output.split("\n"):
|
||||
if line.startswith("diff --git"):
|
||||
# Save previous file
|
||||
if current_file:
|
||||
current_file.diff = "\n".join(current_diff)
|
||||
files.append(current_file)
|
||||
|
||||
# Start new file
|
||||
current_diff = [line]
|
||||
current_file = FileChange(
|
||||
path="",
|
||||
change_type="modified",
|
||||
insertions=0,
|
||||
deletions=0,
|
||||
diff="",
|
||||
)
|
||||
elif line.startswith("--- ") or line.startswith("+++ "):
|
||||
current_diff.append(line)
|
||||
if line.startswith("+++ ") and not line.startswith("+++ /dev/null"):
|
||||
current_file.path = line[6:]
|
||||
elif line.startswith("@@ "):
|
||||
current_diff.append(line)
|
||||
elif line.startswith("+") and not line.startswith("+++"):
|
||||
current_diff.append(line)
|
||||
current_file.insertions += 1
|
||||
elif line.startswith("-") and not line.startswith("---"):
|
||||
current_diff.append(line)
|
||||
current_file.deletions += 1
|
||||
elif current_file:
|
||||
current_diff.append(line)
|
||||
|
||||
# Save last file
|
||||
if current_file:
|
||||
current_file.diff = "\n".join(current_diff)
|
||||
files.append(current_file)
|
||||
|
||||
return files
|
||||
|
||||
|
||||
def _commit_to_dict(commit: Commit, graph_symbol: str = "", graph_depth: int = 0) -> dict[str, Any]:
|
||||
"""Convert Commit dataclass to dictionary."""
|
||||
refs = []
|
||||
if commit.branches:
|
||||
refs.extend(commit.branches)
|
||||
if commit.tags:
|
||||
refs.extend(commit.tags)
|
||||
|
||||
return {
|
||||
"hash": commit.hash,
|
||||
"short_hash": commit.short_hash,
|
||||
"parents": commit.parents,
|
||||
"author_name": commit.author,
|
||||
"author_email": commit.email,
|
||||
"author_date": datetime.fromtimestamp(commit.timestamp, tz=timezone.utc).isoformat(),
|
||||
"message": commit.message,
|
||||
"refs": refs,
|
||||
"graph_symbol": graph_symbol,
|
||||
"graph_depth": graph_depth,
|
||||
}
|
||||
|
||||
|
||||
def _file_change_to_dict(file_change: FileChange) -> dict[str, Any]:
|
||||
"""Convert FileChange dataclass to dictionary."""
|
||||
return {
|
||||
"path": file_change.path,
|
||||
"change_type": file_change.change_type,
|
||||
"insertions": file_change.insertions,
|
||||
"deletions": file_change.deletions,
|
||||
"diff": file_change.diff,
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
"""Git URL parsing utilities to extract base repository URLs from browser URLs."""
|
||||
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
||||
def extract_base_repo_url(url: str) -> str | None:
|
||||
"""Extract base repository URL from a browser/git URL.
|
||||
|
||||
Examples:
|
||||
https://github.com/user/repo/tree/main → https://github.com/user/repo.git
|
||||
https://github.com/user/repo.git → https://github.com/user/repo.git
|
||||
git@github.com:user/repo.git → git@github.com:user/repo.git
|
||||
https://gitlab.com/user/repo/-/blob/main/README.md → https://gitlab.com/user/repo.git
|
||||
|
||||
Returns None if URL doesn't match known patterns.
|
||||
"""
|
||||
# Handle SSH URLs (pass through unchanged)
|
||||
if url.startswith("git@"):
|
||||
return url if url.endswith(".git") else f"{url}.git"
|
||||
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
# Remove query parameters
|
||||
url = f"{parsed.scheme}://{parsed.netloc}{parsed.path}"
|
||||
|
||||
# Extract host
|
||||
host = parsed.netloc.lower()
|
||||
|
||||
# Split path
|
||||
path_parts = [p for p in parsed.path.split("/") if p]
|
||||
|
||||
if not path_parts:
|
||||
return None
|
||||
|
||||
# Determine host type and extract base
|
||||
if "github.com" in host:
|
||||
return _extract_github_url(url, path_parts)
|
||||
elif "gitlab.com" in host:
|
||||
return _extract_gitlab_url(url, path_parts)
|
||||
elif "bitbucket.org" in host:
|
||||
return _extract_bitbucket_url(url, path_parts)
|
||||
else:
|
||||
# Generic host - try basic extraction
|
||||
return _extract_generic_url(url, path_parts)
|
||||
|
||||
|
||||
def _extract_github_url(url: str, path_parts: list[str]) -> str | None:
|
||||
"""Extract base repo URL from GitHub URL."""
|
||||
# Need at least owner/repo
|
||||
if len(path_parts) < 2:
|
||||
return None
|
||||
|
||||
# Find the repo name (second path part)
|
||||
# Remove trailing .git if present
|
||||
repo_name = path_parts[1]
|
||||
if repo_name.endswith(".git"):
|
||||
repo_name = repo_name[:-4]
|
||||
|
||||
# Reconstruct base URL
|
||||
base = f"https://github.com/{path_parts[0]}/{repo_name}"
|
||||
|
||||
# Add .git suffix
|
||||
return f"{base}.git"
|
||||
|
||||
|
||||
def _extract_gitlab_url(url: str, path_parts: list[str]) -> str | None:
|
||||
"""Extract base repo URL from GitLab URL."""
|
||||
# Need at least owner/repo
|
||||
if len(path_parts) < 2:
|
||||
return None
|
||||
|
||||
# Find the repo name (second path part)
|
||||
repo_name = path_parts[1]
|
||||
if repo_name.endswith(".git"):
|
||||
repo_name = repo_name[:-4]
|
||||
|
||||
# Reconstruct base URL
|
||||
base = f"https://gitlab.com/{path_parts[0]}/{repo_name}"
|
||||
|
||||
return f"{base}.git"
|
||||
|
||||
|
||||
def _extract_bitbucket_url(url: str, path_parts: list[str]) -> str | None:
|
||||
"""Extract base repo URL from Bitbucket URL."""
|
||||
# Need at least owner/repo
|
||||
if len(path_parts) < 2:
|
||||
return None
|
||||
|
||||
# Find the repo name (second path part)
|
||||
repo_name = path_parts[1]
|
||||
if repo_name.endswith(".git"):
|
||||
repo_name = repo_name[:-4]
|
||||
|
||||
# Reconstruct base URL
|
||||
base = f"https://bitbucket.org/{path_parts[0]}/{repo_name}"
|
||||
|
||||
return f"{base}.git"
|
||||
|
||||
|
||||
def _extract_generic_url(url: str, path_parts: list[str]) -> str | None:
|
||||
"""Extract base repo URL from generic git host URL."""
|
||||
# Need at least owner/repo
|
||||
if len(path_parts) < 2:
|
||||
return None
|
||||
|
||||
# Find the repo name (second path part)
|
||||
repo_name = path_parts[1]
|
||||
if repo_name.endswith(".git"):
|
||||
repo_name = repo_name[:-4]
|
||||
|
||||
# Reconstruct base URL
|
||||
parsed = urlparse(url)
|
||||
base = f"{parsed.scheme}://{parsed.netloc}/{path_parts[0]}/{repo_name}"
|
||||
|
||||
return f"{base}.git"
|
||||
|
||||
|
||||
def is_valid_clone_url(url: str) -> bool:
|
||||
"""Check if URL is already a valid git clone URL.
|
||||
|
||||
A valid clone URL:
|
||||
- Is an SSH URL (git@host:path)
|
||||
- Ends with .git
|
||||
- Has no browser-specific path segments
|
||||
"""
|
||||
# SSH URLs are always valid
|
||||
if url.startswith("git@"):
|
||||
return True
|
||||
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
path = parsed.path
|
||||
|
||||
# Must end with .git for HTTPS
|
||||
if not path.endswith(".git"):
|
||||
return False
|
||||
|
||||
# Check for browser-specific paths
|
||||
browser_paths = ["/tree/", "/blob/", "/pull/", "/issues/", "/actions/",
|
||||
"-/tree/", "-/blob/", "-/merge_requests/",
|
||||
"/src/"]
|
||||
|
||||
for bp in browser_paths:
|
||||
if bp in path:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def parse_git_url(url: str) -> dict:
|
||||
"""Parse a git URL and return detailed information.
|
||||
|
||||
Returns:
|
||||
{
|
||||
"original_url": str,
|
||||
"base_url": str | None,
|
||||
"is_valid_clone_url": bool,
|
||||
"needs_parsing": bool,
|
||||
"host": str | None,
|
||||
"message": str,
|
||||
"error_code": str | None,
|
||||
}
|
||||
"""
|
||||
result = {
|
||||
"original_url": url,
|
||||
"base_url": None,
|
||||
"is_valid_clone_url": False,
|
||||
"needs_parsing": False,
|
||||
"host": None,
|
||||
"message": "",
|
||||
"error_code": None,
|
||||
}
|
||||
|
||||
# Check if empty
|
||||
if not url or not url.strip():
|
||||
result["message"] = "Please enter a URL"
|
||||
result["error_code"] = "INVALID_URL"
|
||||
return result
|
||||
|
||||
url = url.strip()
|
||||
|
||||
# Try to parse
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except Exception:
|
||||
result["message"] = "Please enter a valid URL"
|
||||
result["error_code"] = "INVALID_URL"
|
||||
return result
|
||||
|
||||
# Extract host
|
||||
if parsed.netloc:
|
||||
result["host"] = parsed.netloc.lower()
|
||||
elif url.startswith("git@"):
|
||||
# SSH URL: git@host:path
|
||||
parts = url.split(":", 1)
|
||||
if len(parts) == 2:
|
||||
result["host"] = parts[0].replace("git@", "")
|
||||
else:
|
||||
result["message"] = "Please enter a valid URL"
|
||||
result["error_code"] = "INVALID_URL"
|
||||
return result
|
||||
|
||||
# Check if already valid
|
||||
if is_valid_clone_url(url):
|
||||
result["base_url"] = url
|
||||
result["is_valid_clone_url"] = True
|
||||
result["needs_parsing"] = False
|
||||
result["message"] = "Valid git repository URL"
|
||||
return result
|
||||
|
||||
# Try to extract base URL
|
||||
base = extract_base_repo_url(url)
|
||||
if base:
|
||||
result["base_url"] = base
|
||||
result["needs_parsing"] = True
|
||||
result["message"] = f"This looks like a browser URL. Did you mean: {base}?"
|
||||
result["error_code"] = "URL_NEEDS_PARSING"
|
||||
else:
|
||||
result["message"] = "Could not parse this URL. Please enter a valid git repository URL."
|
||||
result["error_code"] = "INVALID_URL"
|
||||
|
||||
return result
|
||||
+214
-73
@@ -3,90 +3,231 @@
|
||||
import asyncio
|
||||
import os
|
||||
from typing import AsyncGenerator, Generator
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src.config import Settings, build_database_url
|
||||
# Set test environment BEFORE importing app modules
|
||||
os.environ["APP_ENV"] = "testing"
|
||||
os.environ["SECRET_KEY"] = "test-secret-key-for-testing-only-do-not-use-in-production"
|
||||
os.environ["DATABASE_URL"] = "sqlite+aiosqlite:///:memory:"
|
||||
|
||||
from src.config import Settings
|
||||
from src.models.base import Base
|
||||
from src.main import app
|
||||
|
||||
|
||||
# Unit test fixtures (SQLite in-memory)
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def sqlite_engine():
|
||||
"""Create a SQLite in-memory engine for unit tests."""
|
||||
engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False})
|
||||
Base.metadata.create_all(engine)
|
||||
yield engine
|
||||
engine.dispose()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sqlite_session(sqlite_engine) -> Generator:
|
||||
"""Provide a SQLite session for unit tests."""
|
||||
connection = sqlite_engine.connect()
|
||||
transaction = connection.begin()
|
||||
session = sessionmaker(bind=connection)()
|
||||
|
||||
yield session
|
||||
|
||||
session.close()
|
||||
transaction.rollback()
|
||||
connection.close()
|
||||
|
||||
|
||||
# Integration test fixtures (PostgreSQL)
|
||||
|
||||
TEST_DATABASE_URL = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(scope="session")
|
||||
async def postgres_engine():
|
||||
"""Create a PostgreSQL engine for integration tests."""
|
||||
engine = create_async_engine(TEST_DATABASE_URL)
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
yield engine
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def db_session(postgres_engine) -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Provide a database session with transaction rollback."""
|
||||
async with postgres_engine.connect() as connection:
|
||||
transaction = await connection.begin_nested()
|
||||
session_factory = async_sessionmaker(
|
||||
connection, expire_on_commit=False, class_=AsyncSession
|
||||
)
|
||||
session = session_factory()
|
||||
|
||||
yield session
|
||||
|
||||
await session.close()
|
||||
await transaction.rollback()
|
||||
from src.auth.dependencies import get_db_session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_client() -> Generator[TestClient, None, None]:
|
||||
"""Provide a FastAPI test client."""
|
||||
with TestClient(app) as client:
|
||||
yield client
|
||||
"""Provide a FastAPI test client with SQLite database."""
|
||||
# Create a single engine for this test
|
||||
engine = create_async_engine(
|
||||
"sqlite+aiosqlite:///:memory:",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
# Create tables
|
||||
async def init_db():
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
asyncio.run(init_db())
|
||||
|
||||
async def override_get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with async_sessionmaker(engine, expire_on_commit=False)() as session:
|
||||
yield session
|
||||
|
||||
# Override the dependency
|
||||
app.dependency_overrides[get_db_session] = override_get_db_session
|
||||
|
||||
# Patch startup events to prevent PostgreSQL connection attempts
|
||||
with patch("src.main.init_database") as mock_init:
|
||||
mock_init.return_value = True
|
||||
|
||||
try:
|
||||
with TestClient(app) as client:
|
||||
yield client
|
||||
finally:
|
||||
# Clean up overrides
|
||||
app.dependency_overrides.pop(get_db_session, None)
|
||||
asyncio.run(engine.dispose())
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def configure_test_env(monkeypatch):
|
||||
"""Configure environment for testing."""
|
||||
monkeypatch.setenv("DATABASE_URL", TEST_DATABASE_URL)
|
||||
monkeypatch.setenv("APP_ENV", "testing")
|
||||
@pytest_asyncio.fixture
|
||||
async def db_session(test_client) -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Provide an async database session for unit tests."""
|
||||
# Get the override function from the test_client fixture
|
||||
override_fn = app.dependency_overrides.get(get_db_session)
|
||||
if override_fn:
|
||||
gen = override_fn()
|
||||
session = await gen.asend(None)
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
await gen.aclose()
|
||||
else:
|
||||
# Fallback: create a new engine and session
|
||||
engine = create_async_engine(
|
||||
"sqlite+aiosqlite:///:memory:",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
async with async_sessionmaker(engine, expire_on_commit=False)() as session:
|
||||
yield session
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def authenticated_client(test_client) -> Generator[TestClient, None, None]:
|
||||
"""Provide an authenticated test client with a test user."""
|
||||
import uuid
|
||||
from src.auth.session import create_session_cookie
|
||||
from src.models.user import User
|
||||
|
||||
user_id = str(uuid.uuid4())
|
||||
settings = Settings()
|
||||
|
||||
# Create user in database using the same engine as test_client
|
||||
# We need to access the engine from the test_client fixture
|
||||
# Since we can't easily do that, we'll create the user via API call
|
||||
# But we need the user to exist before any API calls
|
||||
# So we need to create the user using the overridden dependency
|
||||
|
||||
async def create_test_user():
|
||||
# Get the override function
|
||||
override_fn = app.dependency_overrides.get(get_db_session)
|
||||
if override_fn:
|
||||
gen = override_fn()
|
||||
session = await gen.asend(None)
|
||||
try:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email="test@headquarter.local",
|
||||
name="Test User",
|
||||
authentik_id=f"authentik-{user_id}",
|
||||
avatar_url=None,
|
||||
)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
finally:
|
||||
await gen.aclose()
|
||||
|
||||
asyncio.run(create_test_user())
|
||||
|
||||
# Create session cookie
|
||||
session_cookie = create_session_cookie(
|
||||
settings=settings,
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
# Set cookie on client
|
||||
test_client.cookies.set("session", session_cookie)
|
||||
|
||||
yield test_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_project_and_repo(authenticated_client) -> tuple[str, str]:
|
||||
"""Create a project and repository directly in the database."""
|
||||
import uuid
|
||||
from src.models.project import Project
|
||||
from src.models.git_repository import GitRepository
|
||||
|
||||
project_id = uuid.uuid4()
|
||||
repo_id = uuid.uuid4()
|
||||
user_id = None
|
||||
|
||||
# Get user ID from session
|
||||
async def get_user_id():
|
||||
nonlocal user_id
|
||||
from src.auth.session import decode_session_cookie
|
||||
settings = Settings()
|
||||
session_cookie = authenticated_client.cookies.get("session")
|
||||
if session_cookie:
|
||||
session = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||
if session:
|
||||
user_id = uuid.UUID(session["user_id"])
|
||||
|
||||
asyncio.run(get_user_id())
|
||||
|
||||
if not user_id:
|
||||
raise RuntimeError("Could not get user ID from authenticated client")
|
||||
|
||||
async def create_project_and_repo():
|
||||
override_fn = app.dependency_overrides.get(get_db_session)
|
||||
if override_fn:
|
||||
gen = override_fn()
|
||||
session = await gen.asend(None)
|
||||
try:
|
||||
project = Project(
|
||||
id=project_id,
|
||||
name="test-project",
|
||||
description="Test project",
|
||||
owner_id=user_id,
|
||||
)
|
||||
session.add(project)
|
||||
|
||||
repo = GitRepository(
|
||||
id=repo_id,
|
||||
name="test-repo",
|
||||
path="/tmp/test-repo",
|
||||
project_id=project_id,
|
||||
owner_id=user_id,
|
||||
remote_url="https://github.com/test/repo.git",
|
||||
)
|
||||
session.add(repo)
|
||||
await session.commit()
|
||||
finally:
|
||||
await gen.aclose()
|
||||
|
||||
asyncio.run(create_project_and_repo())
|
||||
|
||||
return str(project_id), str(repo_id)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def admin_client(test_client) -> Generator[TestClient, None, None]:
|
||||
"""Provide an authenticated test client with an admin user."""
|
||||
import uuid
|
||||
from src.auth.session import create_session_cookie
|
||||
from src.models.user import User
|
||||
|
||||
user_id = str(uuid.uuid4())
|
||||
settings = Settings()
|
||||
|
||||
async def create_admin_user():
|
||||
override_fn = app.dependency_overrides.get(get_db_session)
|
||||
if override_fn:
|
||||
gen = override_fn()
|
||||
session = await gen.asend(None)
|
||||
try:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email="admin@headquarter.local",
|
||||
name="Admin User",
|
||||
authentik_id=f"authentik-admin-{user_id}",
|
||||
avatar_url=None,
|
||||
is_admin=True,
|
||||
)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
finally:
|
||||
await gen.aclose()
|
||||
|
||||
asyncio.run(create_admin_user())
|
||||
|
||||
# Create session cookie
|
||||
session_cookie = create_session_cookie(
|
||||
settings=settings,
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
# Set cookie on client
|
||||
test_client.cookies.set("session", session_cookie)
|
||||
|
||||
yield test_client
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import asyncio
|
||||
import importlib
|
||||
|
||||
@@ -8,7 +7,7 @@ import pytest
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from src.auth.jwt_service import mint_access_token
|
||||
from src.auth.session import create_session_cookie
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.user import User
|
||||
@@ -27,7 +26,7 @@ def _prepare_auth_test_db() -> None:
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(text("TRUNCATE TABLE refresh_tokens, users RESTART IDENTITY CASCADE"))
|
||||
await connection.execute(text("TRUNCATE TABLE users RESTART IDENTITY CASCADE"))
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
@@ -44,7 +43,7 @@ def _load_app():
|
||||
return main_module.app
|
||||
|
||||
|
||||
def _insert_user_for_refresh(user_id: str) -> None:
|
||||
def _insert_test_user(user_id: str) -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
@@ -63,9 +62,9 @@ def _insert_user_for_refresh(user_id: str) -> None:
|
||||
async with session_factory() as session:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email="refresh@headquarter.local",
|
||||
name="Refresh User",
|
||||
authentik_id="refresh-user",
|
||||
email="test@headquarter.local",
|
||||
name="Test User",
|
||||
authentik_id="test-user",
|
||||
avatar_url=None,
|
||||
)
|
||||
await session.merge(user)
|
||||
@@ -88,7 +87,7 @@ def test_login_redirects_to_authentik_authorize_endpoint() -> None:
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_me_returns_401_without_access_cookie() -> None:
|
||||
def test_me_returns_401_without_session_cookie() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
@@ -99,116 +98,33 @@ def test_me_returns_401_without_access_cookie() -> None:
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_me_returns_user_payload_with_valid_access_cookie() -> None:
|
||||
def test_me_returns_user_with_valid_session() -> None:
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_prepare_auth_test_db()
|
||||
_insert_test_user(user_id)
|
||||
app = _load_app()
|
||||
|
||||
settings = Settings()
|
||||
token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(uuid.uuid4()),
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=15),
|
||||
)
|
||||
session_cookie = create_session_cookie(settings=settings, user_id=user_id)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", token)
|
||||
response = client.get("/auth/me")
|
||||
response = client.get("/auth/me", cookies={"session": session_cookie})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["email"] == "dev@headquarter.local"
|
||||
data = response.json()
|
||||
assert data["email"] == "test@headquarter.local"
|
||||
assert data["name"] == "Test User"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_logout_clears_auth_cookies() -> None:
|
||||
def test_logout_clears_session_cookie() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "opaque-token")
|
||||
response = client.post("/auth/logout")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "access_token=" in response.headers.get("set-cookie", "")
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_callback_rejects_mismatched_state() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("auth_state", "expected")
|
||||
response = client.get("/auth/callback?code=test-code&state=wrong")
|
||||
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_callback_sets_auth_cookies_after_success(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
async def fake_exchange_code_for_tokens(*, settings, code, redirect_uri, client):
|
||||
return {"access_token": "provider-access", "refresh_token": "provider-refresh"}
|
||||
|
||||
def fake_verify_provider_access_token(*, settings, token, jwks):
|
||||
return {"sub": "auth-sub-1", "email": "callback@headquarter.local", "name": "Callback User"}
|
||||
|
||||
async def fake_fetch_jwks(*, settings, client):
|
||||
return {"keys": []}
|
||||
|
||||
monkeypatch.setattr("src.api.auth.exchange_code_for_tokens", fake_exchange_code_for_tokens)
|
||||
monkeypatch.setattr("src.api.auth.verify_provider_access_token", fake_verify_provider_access_token)
|
||||
monkeypatch.setattr("src.api.auth.fetch_jwks", fake_fetch_jwks)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("auth_state", "good-state")
|
||||
response = client.get("/auth/callback?code=valid-code&state=good-state")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["email"] == "callback@headquarter.local"
|
||||
set_cookie_header = response.headers.get("set-cookie", "")
|
||||
assert "access_token=" in set_cookie_header
|
||||
assert "refresh_token=" in set_cookie_header
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_refresh_rotates_cookie_and_returns_user_payload(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
_insert_user_for_refresh("7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb")
|
||||
app = _load_app()
|
||||
|
||||
async def fake_rotate_refresh_token(*, session, raw_token, user_agent, ip_address):
|
||||
class StoredToken:
|
||||
user_id = uuid.UUID("7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb")
|
||||
|
||||
return "new-refresh-token", StoredToken()
|
||||
|
||||
monkeypatch.setattr("src.api.auth.rotate_refresh_token", fake_rotate_refresh_token)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "old-refresh-token")
|
||||
response = client.post("/auth/refresh")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["sub"] == "7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb"
|
||||
assert "refresh_token=" in response.headers.get("set-cookie", "")
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_refresh_returns_401_for_invalid_refresh_token(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
async def fake_rotate_refresh_token(*, session, raw_token, user_agent, ip_address):
|
||||
raise ValueError("refresh token not found")
|
||||
|
||||
monkeypatch.setattr("src.api.auth.rotate_refresh_token", fake_rotate_refresh_token)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "invalid")
|
||||
response = client.post("/auth/refresh")
|
||||
|
||||
assert response.status_code == 401
|
||||
# Check that session cookie is deleted
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert "session=" in set_cookie or "session=\"\"" in set_cookie
|
||||
|
||||
@@ -1,16 +1,9 @@
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import base64
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import build_login_redirect_url, exchange_code_for_tokens, verify_provider_access_token
|
||||
from src.auth.refresh_store import create_refresh_token, hash_refresh_token, revoke_refresh_token, rotate_refresh_token
|
||||
from src.auth.oidc import build_login_redirect_url
|
||||
from src.auth.session import create_session_cookie, decode_session_cookie
|
||||
from src.config import Settings
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -38,153 +31,44 @@ def test_login_redirect_url_contains_required_oidc_params() -> None:
|
||||
settings=settings,
|
||||
redirect_uri="http://localhost:8000/auth/callback",
|
||||
state="state-123",
|
||||
nonce="nonce-123",
|
||||
)
|
||||
|
||||
assert "response_type=code" in url
|
||||
assert "client_id=headquarter-web" in url
|
||||
assert "scope=openid+profile+email" in url
|
||||
assert "state=state-123" in url
|
||||
assert "nonce=nonce-123" in url
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_mint_and_decode_internal_access_token_round_trip() -> None:
|
||||
def test_create_and_decode_session_cookie_round_trip() -> None:
|
||||
settings = Settings()
|
||||
expires_at = datetime.now(UTC) + timedelta(minutes=15)
|
||||
user_id = "test-user-123"
|
||||
|
||||
token = mint_access_token(
|
||||
settings=settings,
|
||||
subject="user-123",
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=expires_at,
|
||||
)
|
||||
cookie = create_session_cookie(settings=settings, user_id=user_id)
|
||||
payload = decode_session_cookie(settings=settings, cookie_value=cookie)
|
||||
|
||||
claims = decode_access_token(settings=settings, token=token)
|
||||
|
||||
assert claims["sub"] == "user-123"
|
||||
assert claims["email"] == "dev@headquarter.local"
|
||||
assert claims["name"] == "Dev User"
|
||||
assert "exp" in claims
|
||||
assert payload["user_id"] == user_id
|
||||
assert "exp" in payload
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_refresh_token_hash_is_deterministic_and_non_reversible() -> None:
|
||||
raw_token = "refresh-token-abc"
|
||||
|
||||
first_hash = hash_refresh_token(raw_token)
|
||||
second_hash = hash_refresh_token(raw_token)
|
||||
|
||||
assert first_hash == second_hash
|
||||
assert first_hash != raw_token
|
||||
assert len(first_hash) == 64
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_decode_access_token_rejects_invalid_signature() -> None:
|
||||
def test_decode_session_rejects_invalid_signature() -> None:
|
||||
settings = Settings()
|
||||
other_settings = Settings(jwt_secret="different-secret")
|
||||
expires_at = datetime.now(UTC) + timedelta(minutes=15)
|
||||
other_settings = Settings(session_secret="different-secret")
|
||||
user_id = "test-user-123"
|
||||
|
||||
token = mint_access_token(
|
||||
settings=other_settings,
|
||||
subject="user-123",
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=expires_at,
|
||||
)
|
||||
cookie = create_session_cookie(settings=other_settings, user_id=user_id)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
decode_access_token(settings=settings, token=token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.integration
|
||||
async def test_exchange_code_for_tokens_posts_expected_payload() -> None:
|
||||
settings = Settings()
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url == httpx.URL(settings.resolved_authentik_token_url)
|
||||
payload = dict(httpx.QueryParams(request.content.decode("utf-8")))
|
||||
assert payload["grant_type"] == "authorization_code"
|
||||
assert payload["code"] == "auth-code"
|
||||
assert payload["redirect_uri"] == "http://localhost:8000/auth/callback"
|
||||
return httpx.Response(200, json={"access_token": "provider-token", "refresh_token": "provider-refresh"})
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(transport=transport) as client:
|
||||
token_payload = await exchange_code_for_tokens(
|
||||
settings=settings,
|
||||
code="auth-code",
|
||||
redirect_uri="http://localhost:8000/auth/callback",
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert token_payload["access_token"] == "provider-token"
|
||||
with pytest.raises(ValueError, match="invalid session signature"):
|
||||
decode_session_cookie(settings=settings, cookie_value=cookie)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_verify_provider_access_token_with_jwks_oct_key() -> None:
|
||||
settings = Settings(authentik_audience="headquarter-web", authentik_issuer="https://authentik.local/")
|
||||
shared_secret = b"shared-secret-123"
|
||||
jwks = {
|
||||
"keys": [
|
||||
{
|
||||
"kty": "oct",
|
||||
"alg": "HS256",
|
||||
"k": base64.urlsafe_b64encode(shared_secret).decode("utf-8").rstrip("="),
|
||||
"kid": "kid-1",
|
||||
}
|
||||
]
|
||||
}
|
||||
def test_decode_session_rejects_expired_cookie(monkeypatch) -> None:
|
||||
settings = Settings(session_ttl_hours=-1) # Already expired
|
||||
user_id = "test-user-123"
|
||||
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
cookie = create_session_cookie(settings=settings, user_id=user_id)
|
||||
|
||||
token = jwt.encode(
|
||||
{
|
||||
"sub": "authentik-user",
|
||||
"iss": settings.authentik_issuer,
|
||||
"aud": settings.authentik_audience,
|
||||
"exp": int((datetime.now(UTC) + timedelta(minutes=5)).timestamp()),
|
||||
},
|
||||
shared_secret,
|
||||
algorithm="HS256",
|
||||
headers={"kid": "kid-1"},
|
||||
)
|
||||
|
||||
claims = verify_provider_access_token(settings=settings, token=token, jwks=jwks)
|
||||
|
||||
assert claims["sub"] == "authentik-user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.integration
|
||||
async def test_refresh_store_create_rotate_and_revoke(db_session: AsyncSession) -> None:
|
||||
user = User(email="dev-auth@headquarter.local", name="Dev Auth", authentik_id="auth-dev", avatar_url=None)
|
||||
db_session.add(user)
|
||||
await db_session.commit()
|
||||
await db_session.refresh(user)
|
||||
|
||||
raw_refresh_token, stored_token = await create_refresh_token(
|
||||
session=db_session,
|
||||
user_id=user.id,
|
||||
expires_at=datetime.now(UTC) + timedelta(days=7),
|
||||
user_agent="pytest",
|
||||
ip_address="127.0.0.1",
|
||||
)
|
||||
assert raw_refresh_token
|
||||
assert stored_token.revoked_at is None
|
||||
|
||||
rotated_raw, rotated_stored = await rotate_refresh_token(
|
||||
session=db_session,
|
||||
raw_token=raw_refresh_token,
|
||||
user_agent="pytest-rotated",
|
||||
ip_address="127.0.0.2",
|
||||
)
|
||||
assert rotated_raw != raw_refresh_token
|
||||
assert rotated_stored.revoked_at is None
|
||||
assert stored_token.revoked_at is not None
|
||||
|
||||
revoked = await revoke_refresh_token(session=db_session, raw_token=rotated_raw)
|
||||
assert revoked is True
|
||||
with pytest.raises(ValueError, match="session expired"):
|
||||
decode_session_cookie(settings=settings, cookie_value=cookie)
|
||||
|
||||
@@ -0,0 +1,255 @@
|
||||
import uuid
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestConfigFoldersAPI:
|
||||
"""Integration tests for config folders API."""
|
||||
|
||||
def test_list_config_folders_requires_authentication(self, test_client: TestClient) -> None:
|
||||
"""Test that listing config folders requires authentication."""
|
||||
response = test_client.get("/config-folders")
|
||||
assert response.status_code == 401
|
||||
|
||||
def test_list_config_folders_returns_user_folders(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that authenticated users can list their folders."""
|
||||
response = authenticated_client.get("/config-folders")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert isinstance(data, dict)
|
||||
assert "folders" in data
|
||||
assert isinstance(data["folders"], list)
|
||||
|
||||
def test_create_config_folder_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a config folder."""
|
||||
response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "test-folder",
|
||||
"description": "Test folder",
|
||||
"mount_path": "/home/user",
|
||||
"files": {"test.txt": "hello world"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "test-folder"
|
||||
assert data["mount_path"] == "/home/user"
|
||||
assert data["files"] == {"test.txt": "hello world"}
|
||||
|
||||
def test_create_config_folder_duplicate_name(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that duplicate folder names are rejected."""
|
||||
# Create first folder
|
||||
response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "duplicate-folder",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
|
||||
# Try to create second with same name
|
||||
response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "duplicate-folder",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 409
|
||||
|
||||
def test_create_config_folder_exceeds_size_limit(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that folders exceeding 10MB are rejected."""
|
||||
large_content = "x" * (11 * 1024 * 1024) # 11MB
|
||||
response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "large-folder",
|
||||
"mount_path": "/home/user",
|
||||
"files": {"large.txt": large_content},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_create_config_folder_path_traversal_attack(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that path traversal in file paths is prevented."""
|
||||
response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "bad-folder",
|
||||
"mount_path": "/home/user",
|
||||
"files": {"../../../etc/passwd": "malicious"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_get_config_folder_by_id(self, authenticated_client: TestClient) -> None:
|
||||
"""Test getting a config folder by ID."""
|
||||
# Create folder first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "get-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
|
||||
# Get it back
|
||||
response = authenticated_client.get(f"/config-folders/{folder_id}")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["name"] == "get-test"
|
||||
|
||||
def test_get_config_folder_not_found(self, authenticated_client: TestClient) -> None:
|
||||
"""Test getting a non-existent folder."""
|
||||
response = authenticated_client.get(f"/config-folders/{uuid.uuid4()}")
|
||||
assert response.status_code == 404
|
||||
|
||||
def test_update_config_folder_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a config folder."""
|
||||
# Create folder first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "update-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
|
||||
# Update it
|
||||
response = authenticated_client.put(
|
||||
f"/config-folders/{folder_id}",
|
||||
json={
|
||||
"name": "updated-name",
|
||||
"mount_path": "/workspace",
|
||||
"files": {"new.txt": "content"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["name"] == "updated-name"
|
||||
assert data["mount_path"] == "/workspace"
|
||||
|
||||
def test_delete_config_folder_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test deleting a config folder."""
|
||||
# Create folder first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "delete-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
|
||||
# Delete it
|
||||
response = authenticated_client.delete(f"/config-folders/{folder_id}")
|
||||
assert response.status_code == 204
|
||||
|
||||
# Verify it's gone
|
||||
get_response = authenticated_client.get(f"/config-folders/{folder_id}")
|
||||
assert get_response.status_code == 404
|
||||
|
||||
def test_add_project_override_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test adding a project override."""
|
||||
# Create folder first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "override-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {"global.txt": "global"},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
project_id = str(uuid.uuid4())
|
||||
|
||||
# Add override
|
||||
response = authenticated_client.post(
|
||||
f"/config-folders/{folder_id}/overrides",
|
||||
json={
|
||||
"project_id": project_id,
|
||||
"mount_path": "/workspace",
|
||||
"files": {"project.txt": "project"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert project_id in data["project_overrides"]
|
||||
|
||||
def test_update_project_override_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a project override."""
|
||||
# Create folder with override
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "update-override-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
project_id = str(uuid.uuid4())
|
||||
|
||||
# Add override
|
||||
authenticated_client.post(
|
||||
f"/config-folders/{folder_id}/overrides",
|
||||
json={
|
||||
"project_id": project_id,
|
||||
"mount_path": "/workspace",
|
||||
"files": {"old.txt": "old"},
|
||||
},
|
||||
)
|
||||
|
||||
# Update override
|
||||
response = authenticated_client.put(
|
||||
f"/config-folders/{folder_id}/overrides/{project_id}",
|
||||
json={
|
||||
"mount_path": "/app",
|
||||
"files": {"new.txt": "new"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["project_overrides"][project_id]["mount_path"] == "/app"
|
||||
|
||||
def test_delete_project_override_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test deleting a project override."""
|
||||
# Create folder with override
|
||||
create_response = authenticated_client.post(
|
||||
"/config-folders",
|
||||
json={
|
||||
"name": "delete-override-test",
|
||||
"mount_path": "/home/user",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
folder_id = create_response.json()["id"]
|
||||
project_id = str(uuid.uuid4())
|
||||
|
||||
# Add override
|
||||
authenticated_client.post(
|
||||
f"/config-folders/{folder_id}/overrides",
|
||||
json={
|
||||
"project_id": project_id,
|
||||
"mount_path": "/workspace",
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
|
||||
# Delete override
|
||||
response = authenticated_client.delete(
|
||||
f"/config-folders/{folder_id}/overrides/{project_id}"
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert project_id not in data["project_overrides"]
|
||||
@@ -0,0 +1,453 @@
|
||||
import uuid
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestConfigProfilesAPI:
|
||||
"""Integration tests for config profiles API."""
|
||||
|
||||
def test_list_config_profiles_requires_authentication(self, test_client: TestClient) -> None:
|
||||
"""Test that listing config profiles requires authentication."""
|
||||
response = test_client.get("/config-profiles")
|
||||
assert response.status_code == 401
|
||||
|
||||
def test_list_config_profiles_returns_user_profiles(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that authenticated users can list their profiles."""
|
||||
response = authenticated_client.get("/config-profiles")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert isinstance(data, list)
|
||||
|
||||
def test_create_config_profile_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a config profile."""
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "test-profile",
|
||||
"description": "Test profile",
|
||||
"env_vars": {"VAR": "value"},
|
||||
"runtime_hints": {"start_command": "npm start"},
|
||||
"mounts": [{"target": "/app", "mode": "rw", "files": {}}],
|
||||
"files": {"test.txt": "hello"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "test-profile"
|
||||
assert data["env_vars"] == {"VAR": "value"}
|
||||
assert data["files"] == {"test.txt": "hello"}
|
||||
assert data["mounts"][0]["target"] == "/app"
|
||||
|
||||
def test_create_config_profile_duplicate_name(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that duplicate profile names are rejected."""
|
||||
# Create first profile
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "duplicate-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
|
||||
# Try to create second with same name
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "duplicate-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 409
|
||||
|
||||
def test_create_config_profile_exceeds_size_limit(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that profiles exceeding 10MB are rejected."""
|
||||
large_content = "x" * (11 * 1024 * 1024) # 11MB
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "large-profile",
|
||||
"env_vars": {},
|
||||
"files": {"large.txt": large_content},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 413
|
||||
|
||||
def test_create_config_profile_invalid_file_path(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that invalid file paths are rejected."""
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "bad-profile",
|
||||
"env_vars": {},
|
||||
"files": {"../../../etc/passwd": "malicious"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_create_config_profile_invalid_mount_target(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that invalid mount targets are rejected."""
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "bad-mount-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"mounts": [{"target": "relative/path", "mode": "rw", "files": {}}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_get_config_profile_by_id(self, authenticated_client: TestClient) -> None:
|
||||
"""Test getting a config profile by ID."""
|
||||
# Create profile first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "get-test",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
profile_id = create_response.json()["id"]
|
||||
|
||||
# Get it back
|
||||
response = authenticated_client.get(f"/config-profiles/{profile_id}")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["name"] == "get-test"
|
||||
|
||||
def test_get_config_profile_not_found(self, authenticated_client: TestClient) -> None:
|
||||
"""Test getting a non-existent profile."""
|
||||
response = authenticated_client.get(f"/config-profiles/{uuid.uuid4()}")
|
||||
assert response.status_code == 404
|
||||
|
||||
def test_update_config_profile_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a config profile."""
|
||||
# Create profile first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "update-test",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
profile_id = create_response.json()["id"]
|
||||
|
||||
# Update it
|
||||
response = authenticated_client.put(
|
||||
f"/config-profiles/{profile_id}",
|
||||
json={
|
||||
"name": "updated-name",
|
||||
"env_vars": {"NEW_VAR": "new_value"},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["name"] == "updated-name"
|
||||
assert data["env_vars"] == {"NEW_VAR": "new_value"}
|
||||
|
||||
def test_delete_config_profile_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test deleting a config profile."""
|
||||
# Create profile first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "delete-test",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
profile_id = create_response.json()["id"]
|
||||
|
||||
# Delete it
|
||||
response = authenticated_client.delete(f"/config-profiles/{profile_id}")
|
||||
assert response.status_code == 204
|
||||
|
||||
# Verify it's gone
|
||||
get_response = authenticated_client.get(f"/config-profiles/{profile_id}")
|
||||
assert get_response.status_code == 404
|
||||
|
||||
def test_update_profile_includes_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating profile includes."""
|
||||
# Create base profile
|
||||
base_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "base-profile",
|
||||
"env_vars": {"BASE_VAR": "base_value"},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
base_id = base_response.json()["id"]
|
||||
|
||||
# Create child profile
|
||||
child_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "child-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
child_id = child_response.json()["id"]
|
||||
|
||||
# Update includes
|
||||
response = authenticated_client.put(
|
||||
f"/config-profiles/{child_id}/includes",
|
||||
json={"includes": [base_id]},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
print(f"Response data: {data}")
|
||||
print(f"Includes: {data.get('includes', 'NO INCLUDES KEY')}")
|
||||
assert len(data["includes"]) == 1, f"Expected 1 include, got {len(data.get('includes', []))}: {data.get('includes', [])}"
|
||||
assert data["includes"][0]["included_profile_id"] == base_id
|
||||
|
||||
def test_update_profile_includes_cycle_detection(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that include cycles are detected."""
|
||||
# Create profile A
|
||||
a_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "profile-a",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
a_id = a_response.json()["id"]
|
||||
|
||||
# Create profile B
|
||||
b_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "profile-b",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
b_id = b_response.json()["id"]
|
||||
|
||||
# Make B include A
|
||||
authenticated_client.put(
|
||||
f"/config-profiles/{b_id}/includes",
|
||||
json={"includes": [a_id]},
|
||||
)
|
||||
|
||||
# Try to make A include B (would create cycle)
|
||||
response = authenticated_client.put(
|
||||
f"/config-profiles/{a_id}/includes",
|
||||
json={"includes": [b_id]},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
|
||||
def test_preview_config_profile_successfully(self, authenticated_client: TestClient) -> None:
|
||||
"""Test previewing a resolved config profile."""
|
||||
# Create base profile
|
||||
base_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "preview-base",
|
||||
"env_vars": {"BASE_VAR": "base"},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
base_id = base_response.json()["id"]
|
||||
|
||||
# Create child profile
|
||||
child_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "preview-child",
|
||||
"env_vars": {"CHILD_VAR": "child"},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
child_id = child_response.json()["id"]
|
||||
|
||||
# Make child include base
|
||||
authenticated_client.put(
|
||||
f"/config-profiles/{child_id}/includes",
|
||||
json={"includes": [base_id]},
|
||||
)
|
||||
|
||||
# Preview child
|
||||
response = authenticated_client.get(f"/config-profiles/{child_id}/preview")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["profile_name"] == "preview-child"
|
||||
assert data["env_vars"]["BASE_VAR"] == "base"
|
||||
assert data["env_vars"]["CHILD_VAR"] == "child"
|
||||
assert len(data["included_profiles"]) == 1
|
||||
|
||||
def test_resolve_default_profile(self, authenticated_client: TestClient) -> None:
|
||||
"""Test resolving default profile for project/tool."""
|
||||
# Create a global default profile (no project/tool scoping)
|
||||
authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "default-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"is_default": True,
|
||||
},
|
||||
)
|
||||
|
||||
# Resolve default with random project/tool (should fall back to global)
|
||||
project_id = str(uuid.uuid4())
|
||||
tool_type_id = str(uuid.uuid4())
|
||||
response = authenticated_client.get(
|
||||
"/config-profiles/defaults/resolve",
|
||||
params={"project_id": project_id, "tool_type_id": tool_type_id},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["profile_name"] == "default-profile"
|
||||
|
||||
def test_resolve_default_profile_no_match(self, authenticated_client: TestClient) -> None:
|
||||
"""Test resolving default profile when no profiles exist."""
|
||||
project_id = str(uuid.uuid4())
|
||||
tool_type_id = str(uuid.uuid4())
|
||||
|
||||
response = authenticated_client.get(
|
||||
"/config-profiles/defaults/resolve",
|
||||
params={"project_id": project_id, "tool_type_id": tool_type_id},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["profile_id"] is None
|
||||
|
||||
def test_create_config_profile_with_git_mounts(self, authenticated_client: TestClient, test_project_and_repo) -> None:
|
||||
"""Test creating a config profile with git mounts."""
|
||||
_project_id, repo_id = test_project_and_repo
|
||||
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "git-mount-profile",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"git_mounts": [
|
||||
{
|
||||
"remote_url": "https://github.com/user/repo.git",
|
||||
"source_path": ".",
|
||||
"target_path": "/app",
|
||||
"branch": "main",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "git-mount-profile"
|
||||
assert len(data["git_mounts"]) == 1
|
||||
assert data["git_mounts"][0]["target_path"] == "/app"
|
||||
assert data["git_mounts"][0]["branch"] == "main"
|
||||
|
||||
def test_update_config_profile_git_mounts(self, authenticated_client: TestClient, test_project_and_repo) -> None:
|
||||
"""Test updating git mounts on a config profile."""
|
||||
_project_id, repo_id = test_project_and_repo
|
||||
|
||||
# Create profile first
|
||||
create_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "update-git-mounts",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
},
|
||||
)
|
||||
profile_id = create_response.json()["id"]
|
||||
|
||||
# Update with git mounts
|
||||
response = authenticated_client.put(
|
||||
f"/config-profiles/{profile_id}",
|
||||
json={
|
||||
"git_mounts": [
|
||||
{
|
||||
"remote_url": "https://github.com/user/repo.git",
|
||||
"source_path": "config",
|
||||
"target_path": "/config",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert len(data["git_mounts"]) == 1
|
||||
assert data["git_mounts"][0]["source_path"] == "config"
|
||||
|
||||
def test_create_config_profile_invalid_git_mount_source_path(self, authenticated_client: TestClient, test_project_and_repo) -> None:
|
||||
"""Test that invalid git mount source paths are rejected."""
|
||||
_project_id, repo_id = test_project_and_repo
|
||||
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "bad-git-mount",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"git_mounts": [
|
||||
{
|
||||
"remote_url": "https://github.com/user/repo.git",
|
||||
"source_path": "/absolute/path",
|
||||
"target_path": "/app",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_create_config_profile_invalid_git_mount_target_path_traversal(self, authenticated_client: TestClient, test_project_and_repo) -> None:
|
||||
"""Test that git mount target paths with traversal are rejected."""
|
||||
_project_id, repo_id = test_project_and_repo
|
||||
|
||||
response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "bad-git-mount-target",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"git_mounts": [
|
||||
{
|
||||
"remote_url": "https://github.com/user/repo.git",
|
||||
"source_path": ".",
|
||||
"target_path": "../../../etc/passwd",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_preview_config_profile_with_git_mounts(self, authenticated_client: TestClient, test_project_and_repo) -> None:
|
||||
"""Test previewing a profile with git mounts."""
|
||||
_project_id, repo_id = test_project_and_repo
|
||||
|
||||
# Create profile with git mounts
|
||||
create_response = authenticated_client.post(
|
||||
"/config-profiles",
|
||||
json={
|
||||
"name": "preview-git-mounts",
|
||||
"env_vars": {},
|
||||
"files": {},
|
||||
"git_mounts": [
|
||||
{
|
||||
"remote_url": "https://github.com/user/repo.git",
|
||||
"source_path": ".",
|
||||
"target_path": "/app",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
profile_id = create_response.json()["id"]
|
||||
|
||||
# Preview
|
||||
response = authenticated_client.get(f"/config-profiles/{profile_id}/preview")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert len(data["git_mounts"]) == 1
|
||||
assert data["git_mounts"][0]["remote_url"] == "https://github.com/user/repo.git"
|
||||
@@ -0,0 +1,164 @@
|
||||
"""Tests for git control utilities."""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
|
||||
from src.utils.git_control import (
|
||||
GitStatus,
|
||||
checkout_branch,
|
||||
commit_changes,
|
||||
create_branch,
|
||||
delete_branch,
|
||||
get_current_branch,
|
||||
get_status,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def temp_repo():
|
||||
"""Create a temporary git repository."""
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
# Initialize git repo
|
||||
os.system(f"cd {tmpdir} && git init && git config user.email 'test@test.com' && git config user.name 'Test User'")
|
||||
|
||||
# Create initial commit
|
||||
with open(os.path.join(tmpdir, "README.md"), "w") as f:
|
||||
f.write("# Test Repo\n")
|
||||
os.system(f"cd {tmpdir} && git add README.md && git commit -m 'Initial commit'")
|
||||
|
||||
yield tmpdir
|
||||
|
||||
|
||||
class TestGitStatus:
|
||||
"""Tests for get_status function."""
|
||||
|
||||
def test_clean_repo(self, temp_repo):
|
||||
"""Test status of a clean repository."""
|
||||
status = get_status(temp_repo)
|
||||
assert isinstance(status, GitStatus)
|
||||
assert status.branch in ["main", "master"]
|
||||
assert len(status.modified) == 0
|
||||
assert len(status.added) == 0
|
||||
assert len(status.deleted) == 0
|
||||
assert len(status.untracked) == 0
|
||||
|
||||
def test_modified_file(self, temp_repo):
|
||||
"""Test detecting modified files."""
|
||||
# Modify a file
|
||||
with open(os.path.join(temp_repo, "README.md"), "w") as f:
|
||||
f.write("# Modified\n")
|
||||
|
||||
status = get_status(temp_repo)
|
||||
assert "README.md" in status.modified
|
||||
|
||||
def test_untracked_file(self, temp_repo):
|
||||
"""Test detecting untracked files."""
|
||||
# Create new file
|
||||
with open(os.path.join(temp_repo, "new.py"), "w") as f:
|
||||
f.write("print('hello')\n")
|
||||
|
||||
status = get_status(temp_repo)
|
||||
assert "new.py" in status.untracked
|
||||
|
||||
|
||||
def test_get_current_branch_handles_unborn_main() -> None:
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
os.system(f"git init -b main {tmpdir} >/dev/null 2>&1")
|
||||
|
||||
assert get_current_branch(tmpdir) == "main"
|
||||
|
||||
|
||||
def test_create_branch_on_bare_repo_with_no_commits() -> None:
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
os.system(f"git init --bare {tmpdir}/bare.git >/dev/null 2>&1")
|
||||
create_branch(f"{tmpdir}/bare.git", "main")
|
||||
assert get_current_branch(f"{tmpdir}/bare.git") == "main"
|
||||
|
||||
|
||||
def test_checkout_branch_on_bare_repo_with_no_commits() -> None:
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
os.system(f"git init --bare {tmpdir}/bare.git >/dev/null 2>&1")
|
||||
checkout_branch(f"{tmpdir}/bare.git", "main")
|
||||
assert get_current_branch(f"{tmpdir}/bare.git") == "main"
|
||||
|
||||
|
||||
class TestBranchOperations:
|
||||
"""Tests for branch management functions."""
|
||||
|
||||
def test_create_branch(self, temp_repo):
|
||||
"""Test creating a new branch."""
|
||||
# Get the actual default branch name
|
||||
default_branch = get_current_branch(temp_repo)
|
||||
create_branch(temp_repo, "feature/test", default_branch)
|
||||
|
||||
# Check branch exists
|
||||
branches = get_status(temp_repo)
|
||||
# Branch should still be on default
|
||||
assert branches.branch == default_branch
|
||||
|
||||
def test_checkout_branch(self, temp_repo):
|
||||
"""Test checking out a branch."""
|
||||
default_branch = get_current_branch(temp_repo)
|
||||
create_branch(temp_repo, "feature/test", default_branch)
|
||||
checkout_branch(temp_repo, "feature/test")
|
||||
|
||||
current = get_current_branch(temp_repo)
|
||||
assert current == "feature/test"
|
||||
|
||||
def test_delete_branch(self, temp_repo):
|
||||
"""Test deleting a branch."""
|
||||
default_branch = get_current_branch(temp_repo)
|
||||
create_branch(temp_repo, "feature/delete", default_branch)
|
||||
delete_branch(temp_repo, "feature/delete")
|
||||
|
||||
# Should be back on default
|
||||
current = get_current_branch(temp_repo)
|
||||
assert current == default_branch
|
||||
|
||||
def test_get_current_branch(self, temp_repo):
|
||||
"""Test getting current branch."""
|
||||
branch = get_current_branch(temp_repo)
|
||||
assert branch in ["main", "master"]
|
||||
|
||||
|
||||
class TestCommit:
|
||||
"""Tests for commit function."""
|
||||
|
||||
def test_commit_changes(self, temp_repo):
|
||||
"""Test committing changes."""
|
||||
# Modify file
|
||||
with open(os.path.join(temp_repo, "README.md"), "w") as f:
|
||||
f.write("# Updated\n")
|
||||
|
||||
# Commit
|
||||
commit_changes(
|
||||
temp_repo,
|
||||
"Update README",
|
||||
"Test User",
|
||||
"test@test.com",
|
||||
["README.md"]
|
||||
)
|
||||
|
||||
# Check status is clean
|
||||
status = get_status(temp_repo)
|
||||
assert "README.md" not in status.modified
|
||||
|
||||
def test_commit_all_changes(self, temp_repo):
|
||||
"""Test committing all changes."""
|
||||
# Modify file
|
||||
with open(os.path.join(temp_repo, "README.md"), "w") as f:
|
||||
f.write("# All updated\n")
|
||||
|
||||
# Commit all
|
||||
commit_changes(
|
||||
temp_repo,
|
||||
"Update all",
|
||||
"Test User",
|
||||
"test@test.com"
|
||||
)
|
||||
|
||||
# Check status is clean
|
||||
status = get_status(temp_repo)
|
||||
assert len(status.modified) == 0
|
||||
@@ -5,7 +5,6 @@ from src.models import Base
|
||||
from src.models.base import TimestampMixin, UUIDPrimaryKeyMixin
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
from src.models.user_config import UserConfig
|
||||
@@ -83,28 +82,6 @@ def test_repository_and_user_config_relationships_are_registered() -> None:
|
||||
assert UserConfig.user.property.mapper.class_ is User
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
|
||||
def test_refresh_token_table_has_required_columns_and_relationships() -> None:
|
||||
columns = RefreshToken.__table__.columns
|
||||
user_fk = next(iter(RefreshToken.__table__.c.user_id.foreign_keys))
|
||||
|
||||
assert set(columns.keys()) == {
|
||||
"id",
|
||||
"user_id",
|
||||
"token_hash",
|
||||
"expires_at",
|
||||
"revoked_at",
|
||||
"user_agent",
|
||||
"ip_address",
|
||||
"created_at",
|
||||
}
|
||||
assert columns["token_hash"].unique is True
|
||||
assert columns["revoked_at"].nullable is True
|
||||
assert user_fk.target_fullname == "users.id"
|
||||
assert RefreshToken.user.property.mapper.class_ is User
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.integration
|
||||
|
||||
|
||||
@@ -3,10 +3,11 @@ from datetime import UTC, datetime, timedelta
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
||||
|
||||
from src.auth.jwt_service import mint_access_token
|
||||
from src.auth.session import create_session_cookie
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.project import Project
|
||||
@@ -53,7 +54,7 @@ def _load_app():
|
||||
|
||||
def _mint_token(user_id: str) -> str:
|
||||
settings = Settings()
|
||||
return mint_access_token(
|
||||
return create_session_cookie(
|
||||
settings=settings,
|
||||
subject=user_id,
|
||||
email="test@headquarter.local",
|
||||
|
||||
@@ -0,0 +1,255 @@
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestToolConfigsAPIExtended:
|
||||
"""Integration tests for tool configs API with new fields."""
|
||||
|
||||
def test_create_tool_config_with_new_fields(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a tool config with all new fields."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "config-test-tool",
|
||||
"display_name": "Config Test Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Create config with new fields
|
||||
response = authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "ADVANCED_CONFIG",
|
||||
"value": "test-value",
|
||||
"config_type": "env",
|
||||
"port_override": 9090,
|
||||
"start_command": "python app.py",
|
||||
"working_directory": "/app",
|
||||
"environment_variables": {"DEBUG": "true", "LOG_LEVEL": "debug"},
|
||||
"volumes": [
|
||||
{"source": "data", "target": "/data", "type": "bind"}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["key"] == "ADVANCED_CONFIG"
|
||||
assert data["port_override"] == 9090
|
||||
assert data["start_command"] == "python app.py"
|
||||
assert data["working_directory"] == "/app"
|
||||
assert data["environment_variables"] == {"DEBUG": "true", "LOG_LEVEL": "debug"}
|
||||
assert data["volumes"] == [{"source": "data", "target": "/data", "type": "bind"}]
|
||||
|
||||
def test_create_tool_config_invalid_port(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that invalid port numbers are rejected."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "port-test-tool",
|
||||
"display_name": "Port Test Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Try to create config with invalid port
|
||||
response = authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "BAD_PORT",
|
||||
"value": "test",
|
||||
"config_type": "env",
|
||||
"port_override": 99999,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_create_tool_config_invalid_volume_structure(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that invalid volume structures are rejected."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "volume-test-tool",
|
||||
"display_name": "Volume Test Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Try to create config with invalid volume
|
||||
response = authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "BAD_VOLUME",
|
||||
"value": "test",
|
||||
"config_type": "env",
|
||||
"volumes": [{"invalid": "structure"}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_update_tool_config_with_new_fields(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a tool config with new fields."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "update-config-tool",
|
||||
"display_name": "Update Config Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Create config
|
||||
create_response = authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "UPDATE_TEST",
|
||||
"value": "original",
|
||||
"config_type": "env",
|
||||
},
|
||||
)
|
||||
config_id = create_response.json()["id"]
|
||||
|
||||
# Update with new fields
|
||||
response = authenticated_client.put(
|
||||
f"/tool-configs/{config_id}",
|
||||
json={
|
||||
"value": "updated",
|
||||
"port_override": 3000,
|
||||
"start_command": "npm start",
|
||||
"working_directory": "/workspace",
|
||||
"environment_variables": {"NODE_ENV": "production"},
|
||||
"volumes": [{"source": "src", "target": "/app/src", "type": "bind"}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["value"] == "updated"
|
||||
assert data["port_override"] == 3000
|
||||
assert data["start_command"] == "npm start"
|
||||
assert data["working_directory"] == "/workspace"
|
||||
assert data["environment_variables"] == {"NODE_ENV": "production"}
|
||||
|
||||
def test_list_tool_configs_returns_new_fields(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that listing configs returns new fields."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "list-config-tool",
|
||||
"display_name": "List Config Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Create config with new fields
|
||||
authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "LIST_TEST",
|
||||
"value": "test",
|
||||
"config_type": "env",
|
||||
"port_override": 5000,
|
||||
"environment_variables": {"TEST": "true"},
|
||||
},
|
||||
)
|
||||
|
||||
# List configs
|
||||
response = authenticated_client.get("/tool-configs")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert len(data) > 0
|
||||
config = data[0]
|
||||
assert "port_override" in config
|
||||
assert "start_command" in config
|
||||
assert "working_directory" in config
|
||||
assert "environment_variables" in config
|
||||
assert "volumes" in config
|
||||
|
||||
def test_get_tool_config_defaults(self, authenticated_client: TestClient) -> None:
|
||||
"""Test getting tool config defaults."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "defaults-tool",
|
||||
"display_name": "Defaults Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx\n volumes:\n - \"{{REPO_PATH}}:/workspace\"\n",
|
||||
"required_variables": ["REPO_PATH"],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Get defaults
|
||||
response = authenticated_client.get(f"/tool-configs/defaults/{tool_id}")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["tool_type_id"] == tool_id
|
||||
assert "suggested_configs" in data
|
||||
|
||||
def test_tool_config_backward_compatibility(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that old configs without new fields still work."""
|
||||
# Create a tool type first
|
||||
tool_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "backward-compat-tool",
|
||||
"display_name": "Backward Compat Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = tool_response.json()["id"]
|
||||
|
||||
# Create config without new fields (simulating old client)
|
||||
response = authenticated_client.post(
|
||||
"/tool-configs",
|
||||
json={
|
||||
"tool_type_id": tool_id,
|
||||
"key": "OLD_STYLE",
|
||||
"value": "value",
|
||||
"config_type": "env",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["key"] == "OLD_STYLE"
|
||||
# New fields should have default values
|
||||
assert data["port_override"] is None
|
||||
assert data["start_command"] is None
|
||||
assert data["working_directory"] is None
|
||||
assert data["environment_variables"] is None
|
||||
assert data["volumes"] is None
|
||||
@@ -0,0 +1,421 @@
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
||||
|
||||
from src.auth.session import create_session_cookie
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.tool_type import ToolType
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
def _prepare_test_db() -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(text("TRUNCATE TABLE tool_types, git_repositories, ssh_keys, projects, users RESTART IDENTITY CASCADE"))
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _load_app():
|
||||
import importlib
|
||||
import src.database as database_module
|
||||
import src.api.auth as auth_module
|
||||
import src.api.tool_types as tool_types_module
|
||||
import src.main as main_module
|
||||
|
||||
# Dispose old engine connections before reload to prevent pool exhaustion
|
||||
if hasattr(database_module, 'engine'):
|
||||
import asyncio
|
||||
asyncio.run(database_module.engine.dispose())
|
||||
|
||||
importlib.reload(database_module)
|
||||
importlib.reload(auth_module)
|
||||
importlib.reload(tool_types_module)
|
||||
importlib.reload(main_module)
|
||||
return main_module.app
|
||||
|
||||
|
||||
def _mint_token(user_id: str) -> str:
|
||||
settings = Settings()
|
||||
return create_session_cookie(
|
||||
settings=settings,
|
||||
subject=user_id,
|
||||
email="test@headquarter.local",
|
||||
name="Test User",
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=15),
|
||||
)
|
||||
|
||||
|
||||
def _insert_user(user_id: str, email: str = "test@headquarter.local") -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory() as session:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email=email,
|
||||
name="Test User",
|
||||
authentik_id=f"authentik-{user_id}",
|
||||
avatar_url=None,
|
||||
)
|
||||
await session.merge(user)
|
||||
await session.commit()
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _insert_tool_type(
|
||||
tool_type_id: str,
|
||||
name: str,
|
||||
display_name: str,
|
||||
compose_template: str,
|
||||
created_by_id: str | None = None,
|
||||
) -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory() as session:
|
||||
tool_type = ToolType(
|
||||
id=uuid.UUID(tool_type_id),
|
||||
name=name,
|
||||
display_name=display_name,
|
||||
description="A test tool type",
|
||||
compose_template=compose_template,
|
||||
required_variables=["REPO_PATH", "TOOL_NAME"],
|
||||
created_by_id=uuid.UUID(created_by_id) if created_by_id else None,
|
||||
)
|
||||
await session.merge(tool_type)
|
||||
await session.commit()
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_list_tool_types_requires_authentication() -> None:
|
||||
_prepare_test_db()
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.get("/tool-types")
|
||||
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_list_tool_types_returns_all_types() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
_insert_tool_type(
|
||||
"aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
|
||||
"custom-tool",
|
||||
"Custom Tool",
|
||||
"version: '3.8'\nservices:\n app:\n image: custom",
|
||||
created_by_id=user_id,
|
||||
)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.get("/tool-types")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert len(data) >= 1
|
||||
custom_tool = next((t for t in data if t["name"] == "custom-tool"), None)
|
||||
assert custom_tool is not None
|
||||
assert custom_tool["display_name"] == "Custom Tool"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_get_tool_type_by_id() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
tool_type_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(user_id)
|
||||
_insert_tool_type(
|
||||
tool_type_id,
|
||||
"custom-tool",
|
||||
"Custom Tool",
|
||||
"version: '3.8'\nservices:\n app:\n image: custom",
|
||||
created_by_id=user_id,
|
||||
)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.get(f"/tool-types/{tool_type_id}")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["id"] == tool_type_id
|
||||
assert data["name"] == "custom-tool"
|
||||
assert data["display_name"] == "Custom Tool"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_get_tool_type_not_found() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.get("/tool-types/aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")
|
||||
|
||||
assert response.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_create_tool_type_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"name": "my-custom-tool",
|
||||
"display_name": "My Custom Tool",
|
||||
"description": "A custom development tool",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: custom:latest",
|
||||
"required_variables": ["REPO_PATH"],
|
||||
}
|
||||
response = client.post("/tool-types", json=payload)
|
||||
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "my-custom-tool"
|
||||
assert data["display_name"] == "My Custom Tool"
|
||||
assert data["description"] == "A custom development tool"
|
||||
assert data["created_by_id"] == user_id
|
||||
assert "id" in data
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_create_tool_type_duplicate_name() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
_insert_tool_type(
|
||||
"aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
|
||||
"existing-tool",
|
||||
"Existing Tool",
|
||||
"version: '3.8'\nservices:\n app:\n image: existing",
|
||||
created_by_id=user_id,
|
||||
)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"name": "existing-tool",
|
||||
"display_name": "Existing Tool",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: custom",
|
||||
"required_variables": [],
|
||||
}
|
||||
response = client.post("/tool-types", json=payload)
|
||||
|
||||
assert response.status_code == 409
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_create_tool_type_invalid_yaml() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"name": "bad-tool",
|
||||
"display_name": "Bad Tool",
|
||||
"compose_template": "this is not: valid: yaml: [",
|
||||
"required_variables": [],
|
||||
}
|
||||
response = client.post("/tool-types", json=payload)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_create_tool_type_missing_services() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"name": "bad-tool",
|
||||
"display_name": "Bad Tool",
|
||||
"compose_template": "version: '3.8'\ninvalid_key: value",
|
||||
"required_variables": [],
|
||||
}
|
||||
response = client.post("/tool-types", json=payload)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_create_tool_type_missing_required_variable() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"name": "bad-tool",
|
||||
"display_name": "Bad Tool",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: custom",
|
||||
"required_variables": ["MISSING_VAR"],
|
||||
}
|
||||
response = client.post("/tool-types", json=payload)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_update_tool_type_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
tool_type_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(user_id)
|
||||
_insert_tool_type(
|
||||
tool_type_id,
|
||||
"custom-tool",
|
||||
"Custom Tool",
|
||||
"version: '3.8'\nservices:\n app:\n image: custom",
|
||||
created_by_id=user_id,
|
||||
)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {
|
||||
"display_name": "Updated Custom Tool",
|
||||
"description": "Updated description",
|
||||
}
|
||||
response = client.put(f"/tool-types/{tool_type_id}", json=payload)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["display_name"] == "Updated Custom Tool"
|
||||
assert data["description"] == "Updated description"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_update_tool_type_not_found() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
payload = {"display_name": "Updated"}
|
||||
response = client.put("/tool-types/aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", json=payload)
|
||||
|
||||
assert response.status_code == 404
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_delete_tool_type_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
tool_type_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(user_id)
|
||||
_insert_tool_type(
|
||||
tool_type_id,
|
||||
"deletable-tool",
|
||||
"Deletable Tool",
|
||||
"version: '3.8'\nservices:\n app:\n image: custom",
|
||||
created_by_id=user_id,
|
||||
)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.delete(f"/tool-types/{tool_type_id}")
|
||||
|
||||
assert response.status_code == 204
|
||||
|
||||
# Verify it's gone
|
||||
get_response = client.get(f"/tool-types/{tool_type_id}")
|
||||
assert get_response.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_delete_tool_type_not_found() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.delete("/tool-types/aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")
|
||||
|
||||
assert response.status_code == 404
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestToolTypesAPIExtended:
|
||||
"""Integration tests for tool types API with new fields."""
|
||||
|
||||
def test_create_tool_type_with_dockerfile(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a tool type with dockerfile definition."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "dockerfile-tool",
|
||||
"display_name": "Dockerfile Tool",
|
||||
"category": "utility",
|
||||
"interfaces": ["terminal"],
|
||||
"default_port": 8080,
|
||||
"definition_type": "dockerfile",
|
||||
"dockerfile_template": "FROM python:3.11\nRUN pip install flask",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "dockerfile-tool"
|
||||
assert data["definition_type"] == "dockerfile"
|
||||
assert data["dockerfile_template"] == "FROM python:3.11\nRUN pip install flask"
|
||||
|
||||
def test_create_tool_type_with_readiness_probe(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a tool type with readiness probe."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "probed-tool",
|
||||
"display_name": "Probed Tool",
|
||||
"category": "utility",
|
||||
"interfaces": ["web"],
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx\n ports:\n - '8080:8080'",
|
||||
"readiness_probe": {
|
||||
"command": "curl -f http://localhost:8080",
|
||||
"timeout": 30,
|
||||
"interval": 2,
|
||||
},
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["readiness_probe"]["command"] == "curl -f http://localhost:8080"
|
||||
assert data["readiness_probe"]["timeout"] == 30
|
||||
|
||||
def test_create_tool_type_invalid_definition_type(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that invalid definition types are rejected."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "invalid-tool",
|
||||
"display_name": "Invalid Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "invalid",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_create_tool_type_dockerfile_without_template(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that dockerfile type requires dockerfile_template."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "no-dockerfile",
|
||||
"display_name": "No Dockerfile",
|
||||
"default_port": 8080,
|
||||
"definition_type": "dockerfile",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_update_tool_type_with_new_fields(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a tool type with new fields."""
|
||||
# Create tool type first
|
||||
create_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "update-test-tool",
|
||||
"display_name": "Update Test Tool",
|
||||
"default_port": 8080,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx\n ports:\n - '8080:8080'",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = create_response.json()["id"]
|
||||
|
||||
# Update it
|
||||
response = authenticated_client.put(
|
||||
f"/tool-types/{tool_id}",
|
||||
json={
|
||||
"display_name": "Updated Name",
|
||||
"readiness_probe": {
|
||||
"command": "curl -f http://localhost:8080/health",
|
||||
"timeout": 60,
|
||||
"interval": 5,
|
||||
},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["display_name"] == "Updated Name"
|
||||
assert data["readiness_probe"]["command"] == "curl -f http://localhost:8080/health"
|
||||
|
||||
def test_validate_tool_type_compose(self, authenticated_client: TestClient) -> None:
|
||||
"""Test validating compose template."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types/validate",
|
||||
json={
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["valid"] is True
|
||||
|
||||
def test_validate_tool_type_invalid_compose(self, authenticated_client: TestClient) -> None:
|
||||
"""Test validating invalid compose template."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types/validate",
|
||||
json={
|
||||
"definition_type": "compose",
|
||||
"compose_template": "invalid: yaml: [",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["valid"] is False
|
||||
assert "errors" in data
|
||||
|
||||
def test_validate_tool_type_dockerfile(self, authenticated_client: TestClient) -> None:
|
||||
"""Test validating dockerfile template."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types/validate",
|
||||
json={
|
||||
"definition_type": "dockerfile",
|
||||
"dockerfile_template": "FROM python:3.11\nRUN pip install flask",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["valid"] is True
|
||||
|
||||
def test_get_tool_type_returns_new_fields(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that GET returns new fields."""
|
||||
# Create tool type with all fields
|
||||
create_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "full-tool",
|
||||
"display_name": "Full Tool",
|
||||
"category": "editor",
|
||||
"interfaces": ["web", "terminal"],
|
||||
"default_port": 8443,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: code-server\n ports:\n - '8443:8443'\n volumes:\n - \"{{REPO_PATH}}:/workspace\"",
|
||||
"readiness_probe": {
|
||||
"command": "curl -f http://localhost:8443",
|
||||
"timeout": 30,
|
||||
"interval": 2,
|
||||
},
|
||||
"required_variables": ["REPO_PATH"],
|
||||
},
|
||||
)
|
||||
tool_id = create_response.json()["id"]
|
||||
|
||||
# Get it
|
||||
response = authenticated_client.get(f"/tool-types/{tool_id}")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["definition_type"] == "compose"
|
||||
assert data["category"] == "editor"
|
||||
assert data["interfaces"] == ["web", "terminal"]
|
||||
assert "readiness_probe" in data
|
||||
|
||||
def test_create_tool_type_without_port_fails(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that creating a tool type without default_port fails validation."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "no-port-tool",
|
||||
"display_name": "No Port Tool",
|
||||
"category": "utility",
|
||||
"interfaces": ["web"],
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx\n ports:\n - '8080:8080'",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
data = response.json()
|
||||
assert "default_port" in str(data)
|
||||
|
||||
def test_create_tool_type_with_port_mismatch_fails(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that port mismatch between default_port and compose template fails."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "port-mismatch-tool",
|
||||
"display_name": "Port Mismatch Tool",
|
||||
"category": "utility",
|
||||
"interfaces": ["web"],
|
||||
"default_port": 9999,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: nginx\n ports:\n - '8080:8080'",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
_ = response.json()
|
||||
|
||||
def test_create_tool_type_with_startup_command(self, authenticated_client: TestClient) -> None:
|
||||
"""Test creating a tool type with startup_command."""
|
||||
response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "startup-tool",
|
||||
"display_name": "Startup Tool",
|
||||
"category": "utility",
|
||||
"interface_type": "terminal",
|
||||
"requires_port": False,
|
||||
"default_port": 0,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: alpine",
|
||||
"startup_command": "cd /workspace && ls",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["startup_command"] == "cd /workspace && ls"
|
||||
assert data["interface_type"] == "terminal"
|
||||
|
||||
def test_update_tool_type_startup_command(self, authenticated_client: TestClient) -> None:
|
||||
"""Test updating a tool type's startup_command."""
|
||||
# Create tool type first
|
||||
create_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "update-startup-tool",
|
||||
"display_name": "Update Startup Tool",
|
||||
"interface_type": "terminal",
|
||||
"requires_port": False,
|
||||
"default_port": 0,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: alpine",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = create_response.json()["id"]
|
||||
|
||||
# Update with startup_command
|
||||
response = authenticated_client.put(
|
||||
f"/tool-types/{tool_id}",
|
||||
json={
|
||||
"startup_command": "source /etc/profile",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["startup_command"] == "source /etc/profile"
|
||||
|
||||
def test_get_tool_type_returns_startup_command(self, authenticated_client: TestClient) -> None:
|
||||
"""Test that GET returns startup_command."""
|
||||
create_response = authenticated_client.post(
|
||||
"/tool-types",
|
||||
json={
|
||||
"name": "get-startup-tool",
|
||||
"display_name": "Get Startup Tool",
|
||||
"interface_type": "terminal",
|
||||
"requires_port": False,
|
||||
"default_port": 0,
|
||||
"definition_type": "compose",
|
||||
"compose_template": "version: '3.8'\nservices:\n app:\n image: alpine",
|
||||
"startup_command": "echo hello",
|
||||
"required_variables": [],
|
||||
},
|
||||
)
|
||||
tool_id = create_response.json()["id"]
|
||||
|
||||
response = authenticated_client.get(f"/tool-types/{tool_id}")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["startup_command"] == "echo hello"
|
||||
assert "Port 9999 is not exposed" in str(data)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user